From d10009d639d4cc68b288fc7337f79b7888ce6042 Mon Sep 17 00:00:00 2001 From: chaos Date: Thu, 16 Jul 2026 21:04:05 +0800 Subject: [PATCH] Initial commit: grok-free-register-oss Open-source Grok free registration CLI, xai_enroller auth pipeline, local auth service, tests and docs. --- .env.example | 124 ++ .gitignore | 19 + LICENSE | 21 + README.md | 370 ++++ auth-service.sh | 55 + cloudflare/email-worker.js | 61 + docs/architecture.md | 180 ++ docs/guides/auth-service.md | 67 + docs/guides/credential-inventory.md | 19 + docs/guides/registration.md | 54 + docs/guides/runtime-troubleshooting.md | 38 + grok_register/__init__.py | 1 + grok_register/clearance.py | 395 ++++ grok_register/core/__init__.py | 0 grok_register/core/admission.py | 129 ++ grok_register/core/envelope.py | 61 + grok_register/core/inventory.py | 169 ++ grok_register/core/observer.py | 215 ++ grok_register/email_server.py | 171 ++ grok_register/register.py | 2308 ++++++++++++++++++++ keys/acpa_watchdog.py | 944 ++++++++ keys/sync_acc.py | 188 ++ requirements.txt | 7 + reset_pipeline.sh | 222 ++ scripts/clean_history.sh | 140 ++ scripts/ensure_runtime.sh | 32 + scripts/export_registered_sessions.py | 206 ++ scripts/push_keys_to_auth.sh | 153 ++ setup.sh | 50 + start.sh | 80 + tests/__init__.py | 0 tests/conftest.py | 92 + tests/fakes.py | 317 +++ tests/requirements.txt | 2 + tests/test_admission_gate.py | 82 + tests/test_cancel.py | 475 ++++ tests/test_core.py | 491 +++++ tests/test_inventory_unittest.py | 58 + tests/test_property.py | 482 ++++ tests/test_register_runtime_unittest.py | 1325 +++++++++++ tests/test_runtime_log_analyzer.py | 132 ++ tests/test_shell_entrypoints.py | 64 + tests/test_stress.py | 411 ++++ tests/test_xai_auth_pipeline.py | 298 +++ tests/test_xai_auth_service.py | 391 ++++ tests/test_xai_credential_inventory.py | 58 + tests/test_xai_enroller_config.py | 125 ++ tests/test_xai_enroller_coordinator.py | 293 +++ tests/test_xai_enroller_executors.py | 411 ++++ tests/test_xai_enroller_integration.py | 188 ++ tests/test_xai_enroller_ledger.py | 163 ++ tests/test_xai_enroller_protocol.py | 245 +++ tests/test_xai_enroller_sinks.py | 79 + tests/test_xai_remote_snapshot.py | 240 ++ tools/GrokSession2CPAAndSub2API/README.md | 100 + tools/GrokSession2CPAAndSub2API/index.html | 902 ++++++++ tools/__init__.py | 1 + tools/runtime_log_analyzer.py | 419 ++++ xai_enroller/__init__.py | 3 + xai_enroller/__main__.py | 111 + xai_enroller/auth_pipeline.py | 1035 +++++++++ xai_enroller/config.py | 139 ++ xai_enroller/coordinator.py | 177 ++ xai_enroller/executors.py | 595 +++++ xai_enroller/inventory.py | 105 + xai_enroller/ledger.py | 275 +++ xai_enroller/models.py | 128 ++ xai_enroller/protocol.py | 181 ++ xai_enroller/remote_stream.py | 630 ++++++ xai_enroller/service.py | 883 ++++++++ xai_enroller/sinks.py | 106 + xai_enroller/sources.py | 61 + 72 files changed, 18752 insertions(+) create mode 100644 .env.example create mode 100644 .gitignore create mode 100644 LICENSE create mode 100644 README.md create mode 100755 auth-service.sh create mode 100644 cloudflare/email-worker.js create mode 100644 docs/architecture.md create mode 100644 docs/guides/auth-service.md create mode 100644 docs/guides/credential-inventory.md create mode 100644 docs/guides/registration.md create mode 100644 docs/guides/runtime-troubleshooting.md create mode 100644 grok_register/__init__.py create mode 100644 grok_register/clearance.py create mode 100644 grok_register/core/__init__.py create mode 100644 grok_register/core/admission.py create mode 100644 grok_register/core/envelope.py create mode 100644 grok_register/core/inventory.py create mode 100644 grok_register/core/observer.py create mode 100644 grok_register/email_server.py create mode 100644 grok_register/register.py create mode 100755 keys/acpa_watchdog.py create mode 100755 keys/sync_acc.py create mode 100644 requirements.txt create mode 100755 reset_pipeline.sh create mode 100755 scripts/clean_history.sh create mode 100755 scripts/ensure_runtime.sh create mode 100755 scripts/export_registered_sessions.py create mode 100755 scripts/push_keys_to_auth.sh create mode 100755 setup.sh create mode 100755 start.sh create mode 100644 tests/__init__.py create mode 100644 tests/conftest.py create mode 100644 tests/fakes.py create mode 100644 tests/requirements.txt create mode 100644 tests/test_admission_gate.py create mode 100644 tests/test_cancel.py create mode 100644 tests/test_core.py create mode 100644 tests/test_inventory_unittest.py create mode 100644 tests/test_property.py create mode 100644 tests/test_register_runtime_unittest.py create mode 100644 tests/test_runtime_log_analyzer.py create mode 100644 tests/test_shell_entrypoints.py create mode 100644 tests/test_stress.py create mode 100644 tests/test_xai_auth_pipeline.py create mode 100644 tests/test_xai_auth_service.py create mode 100644 tests/test_xai_credential_inventory.py create mode 100644 tests/test_xai_enroller_config.py create mode 100644 tests/test_xai_enroller_coordinator.py create mode 100644 tests/test_xai_enroller_executors.py create mode 100644 tests/test_xai_enroller_integration.py create mode 100644 tests/test_xai_enroller_ledger.py create mode 100644 tests/test_xai_enroller_protocol.py create mode 100644 tests/test_xai_enroller_sinks.py create mode 100644 tests/test_xai_remote_snapshot.py create mode 100644 tools/GrokSession2CPAAndSub2API/README.md create mode 100644 tools/GrokSession2CPAAndSub2API/index.html create mode 100644 tools/__init__.py create mode 100644 tools/runtime_log_analyzer.py create mode 100644 xai_enroller/__init__.py create mode 100644 xai_enroller/__main__.py create mode 100644 xai_enroller/auth_pipeline.py create mode 100644 xai_enroller/config.py create mode 100644 xai_enroller/coordinator.py create mode 100644 xai_enroller/executors.py create mode 100644 xai_enroller/inventory.py create mode 100644 xai_enroller/ledger.py create mode 100644 xai_enroller/models.py create mode 100644 xai_enroller/protocol.py create mode 100644 xai_enroller/remote_stream.py create mode 100644 xai_enroller/service.py create mode 100644 xai_enroller/sinks.py create mode 100644 xai_enroller/sources.py diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..0366f46 --- /dev/null +++ b/.env.example @@ -0,0 +1,124 @@ +# grok-free-register 配置模板 +# 复制为 .env 后按需修改。未填写的项目会使用程序默认值。 + +# 邮箱模式 +# tempmail: 免费临时邮箱,适合快速试跑。 +# custom: 自建域名邮箱,需要 Cloudflare Email Routing 和本地收信服务。 +EMAIL_MODE=tempmail + +# custom 模式专用 +# EMAIL_DOMAIN=example.com +# EMAIL_API=http://127.0.0.1:8080 + +# 认证账号来源 +# auto: 配置 SSH 主机时读取远端,否则读取当前项目;一般保持默认即可。 +# XAI_AUTH_SERVICE_SOURCE=auto # auto | local | ssh +# XAI_AUTH_SERVICE_REGISTER_ROOT=. # local 模式下的注册项目目录 + +# 注册与认证不在同一台机器时填写 +# XAI_AUTH_SERVICE_SSH_HOST=user@server.example +# XAI_AUTH_SERVICE_SSH_IDENTITY=/path/to/key.pem +# XAI_AUTH_SERVICE_REMOTE_ROOT=/opt/grok-free-register +# XAI_AUTH_SERVICE_SYNC_SEC=30 # 同步新账号的间隔(秒) + +# 认证节奏 +# XAI_AUTH_SERVICE_MIN_INTERVAL_SEC=10 # 两次认证请求的基础最小间隔(秒) +# # 触发限流后会自动抬高(约18s起,最多约90s),成功连串后再缓慢回落 +# XAI_AUTH_SERVICE_RETRY_SEC=60 # 被限流后的基础探测间隔(秒);连续撞限会按 1.5x 增长,上限 300s +# XAI_AUTH_SERVICE_LOG_MODE=user # user | debug + +# 运行目标 +# TARGET=0 # 成功 N 个后停止;0=不限 + +# 日志显示 +# REGISTER_LOG_MODE=user # 默认:任务开始、结果、平均速度和限流等待 +# REGISTER_LOG_MODE=debug # 额外显示 T/Q/并发/阶段耗时调试面板 + +# 浏览器并发容量 +# PHYSICAL_CAP=0 # 0=按 CPU/内存自动估算;>0=手动指定 +# PHYSICAL_PER_CPU=2 # 自动估算时每个 CPU 核心对应的并发参考值 +# PHYSICAL_MEM_MB=512 # 自动估算时每个浏览器任务的内存预算(MB) +# MIN_FREE_MEM_MB=500 # 自动估算时保留的内存(MB) +# CAPACITY_PROFILE= # 可选:离线压测生成的静态 profile(JSON) + +# 缓冲容量 +# T_SLOT_CAP=8 +# Q_SLOT_CAP=8 +# Q_PENDING_CAP=12 + +# Worker 数量 +# 通常不需要手动设置。0 表示根据容量自动派生。 +# S_WORKERS=0 +# P_WORKERS=0 +# C_WORKERS=0 + +# 资源有效期 +# T_MAX_AGE=300 # token 最大年龄(秒) +# Q_MAX_AGE=120 # 验证码最大年龄(秒) + +# 阶段超时 +# P_REQUEST_TIMEOUT=95 # 等待验证码返回超时(秒) +# C_CONSUME_TIMEOUT=60 # 最终提交阶段超时(秒) + +# 局部门控和批量发送 +# 这些项目默认会按容量派生。只有在压测定位后才建议覆盖。 +# T_TARGET=4 # token 缓冲目标 +# Q_TARGET=4 # 验证码缓冲目标 +# T_HIGH_WATER= # token 生产高水位 +# T_LOW_WATER= # token 生产恢复低水位 +# Q_HIGH_WATER= # 验证码生产高水位 +# Q_LOW_WATER= # 验证码生产恢复低水位 +# P_BATCH_MAX=4 # 单个发送页面最多发送的验证码请求数 +# P_SEND_CAP=0 # 0=不额外限制;>0=限制发送页面并发 + +# 页面和 token 获取 +# SOLVER_REUSE=1 # token 页面复用 +# MAX_SOLVER_REUSE=25 # 单个 token 页面最大复用次数 +# SOLVER_INITIAL_WAIT_MS=500 # token 注入后的首次等待 +# SOLVER_POLL_INTERVAL_MS=500 # token 轮询间隔 +# SOLVER_POLL_ATTEMPTS=100 # token 最大轮询次数 +# SOLVER_HARD_TIMEOUT=90 # 单次 solver 全链路硬截止(秒) +# SOLVER_CLEANUP_TIMEOUT=5 # 异常页回收最长等待(秒) +# SOLVER_FAST_CLICK=1 # 没有可见验证框时跳过慢点击 +# SOLVER_MOUSE_CLICK_RETRIES=3 # token 验证框中心点击次数;0=关闭 +# SOLVER_MOUSE_CLICK_INTERVAL_MS=600 # token 中心点击间隔(ms) +# PAGE_GOTO_WAIT_UNTIL=domcontentloaded # 页面导航等待策略 +# PAGE_POST_WAIT_MS=500 # 页面导航后的短等待 + +# 可选性能开关 +# PAGE_BLOCK_STATIC_ASSETS=0 # 阻断部分静态资源,降低页面准备成本 +# C_HOT_PAGE_POOL=0 # 复用消费阶段页面,会增加常驻内存 +# C_HOT_PAGE_POOL_SIZE=0 # 0=按启动期并发派生;>0=手动指定 +# C_SET_COOKIE_VIA_REQUEST=0 # 用浏览器 request 访问 set-cookie URL;热页池开启时默认启用 + +# --------------------------------------------------------------------------- +# 代理(grok-free-register 本进程出口,给 Playwright + httpx) +# 与下面 FlareSolverr 清障栈是两套东西,不要混。 +# Playwright 不会自动读环境变量,代码里会显式挂到 launch(proxy=...)。 +# --------------------------------------------------------------------------- +# REGISTER_PROXY=http://127.0.0.1:40080 # 优先;指向本机 privoxy→WARP 时用这个 +# HTTP_PROXY=http://127.0.0.1:7890 # 兜底(REGISTER_PROXY 未设时) +# HTTPS_PROXY=http://127.0.0.1:7890 + +# --------------------------------------------------------------------------- +# Cloudflare 清障预热(调用外部 FlareSolverr,不在本仓起 FS) +# 默认 FS: docker compose 暴露 127.0.0.1:8191 +# 预热根域名(无 path 后缀): accounts.x.ai / x.ai / status.x.ai / console.x.ai / auth.x.ai +# CLEARANCE_PROXY 是 FS *容器内* 浏览器的出口,必须是容器可达地址; +# 不要把 host 的 127.0.0.1 塞进去。走 WARP 时用 docker 网内 privoxy。 +# --------------------------------------------------------------------------- +# CLEARANCE_ENABLED=1 +# FLARESOLVERR_URL=http://127.0.0.1:8191 +# CLEARANCE_URLS=https://accounts.x.ai,https://x.ai,https://status.x.ai,https://console.x.ai,https://auth.x.ai +# CLEARANCE_PROXY=http://privoxy:8118 +# CLEARANCE_TIMEOUT_SEC=60 +# CLEARANCE_REFRESH_SEC=3000 + +# --------------------------------------------------------------------------- +# 本机注册 → 推 keys 到远端 auth(可选;scripts/push_keys_to_auth.sh) +# 无内置默认主机,必须显式填写。 +# --------------------------------------------------------------------------- +# AUTH_SSH_HOST=user@auth-server.example +# AUTH_REMOTE_ROOT=/opt/grok-free-register +# AUTH_SSH_IDENTITY= +# AUTH_PUSH_INTERVAL=20 diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..906af29 --- /dev/null +++ b/.gitignore @@ -0,0 +1,19 @@ +__pycache__/ +*.pyc +.env +*.log +logs/ +*.bak +.venv/ +.setup.lock/ +.DS_Store +.pytest_cache/ +xai-enroller-ledger.db* +tmp_*.py + +# runtime credential outputs under keys/ (keep product scripts tracked) +keys/* +!keys/acpa_watchdog.py +!keys/sync_acc.py +!keys/async_auth.sh +!keys/README.md diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..90b7071 --- /dev/null +++ b/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 grok-free-register contributors + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md new file mode 100644 index 0000000..d3d1c92 --- /dev/null +++ b/README.md @@ -0,0 +1,370 @@ +# grok-free-register + +`grok-free-register` 是一个命令行注册工具。程序会启动本机浏览器,完成页面操作、邮箱验证码处理和结果保存。 + +运行结果写入 `keys/` 目录。 + +## 快速开始 + +```bash +git clone /grok-free-register.git +cd grok-free-register +bash start.sh +``` + +首次运行会自动创建 `.venv`、安装依赖,并引导生成 `.env`。 + +需要完整说明时,按用途查看: + +- [注册教程](docs/guides/registration.md) +- [本地认证服务](docs/guides/auth-service.md) +- [凭据库存与取用](docs/guides/credential-inventory.md) +- [运行状态与排障](docs/guides/runtime-troubleshooting.md) + +常用命令: + +```bash +bash start.sh # 按当前 .env 前台运行 +bash start.sh --target 100 # 成功 100 个后停止 +bash start.sh --max-mem 6G # 自动估算并发时最多使用 6G 内存 +bash start.sh --reconfig # 重新选择邮箱模式 +``` + +`start.sh` 直接在当前终端显示状态。按 `Ctrl-C` 停止,再次执行同一命令即可重启;不需要额外的会话管理或守护进程依赖。 + +代理与 CF 清障是两套配置,不要混: + +```env +# 本进程出口(Playwright + httpx 显式挂载;Playwright 不会自己读 env) +REGISTER_PROXY=http://127.0.0.1:40080 +# 或兜底: +# HTTP_PROXY=http://127.0.0.1:7890 +# HTTPS_PROXY=http://127.0.0.1:7890 + +# 清障预热:调用外部 FlareSolverr(本仓不跑 FS) +CLEARANCE_ENABLED=1 +FLARESOLVERR_URL=http://127.0.0.1:8191 +CLEARANCE_URLS=https://accounts.x.ai,https://x.ai,https://status.x.ai,https://console.x.ai,https://auth.x.ai +# FS 容器内出口(容器可达名);走 WARP 时用 docker 网内 privoxy +# CLEARANCE_PROXY=http://privoxy:8118 +``` + +## 邮箱模式 + +`tempmail` 是默认模式,不需要额外配置,适合快速试跑: + +```env +EMAIL_MODE=tempmail +``` + +`custom` 是自建域名邮箱模式,适合长时间运行。需要一个已接入 Cloudflare Email Routing 的域名,并在运行机器上启动本项目的收信服务。 + +配置步骤: + +1. 在 Cloudflare 为域名开启 Email Routing。 +2. 部署 `cloudflare/email-worker.js`。 +3. 在 Email Routing 中配置 catch-all,动作选择发送到该 Worker。 +4. 在运行机器上启动收信服务: + +```bash +bash start.sh --email-service +``` + +5. 在 `.env` 中配置: + +```env +EMAIL_MODE=custom +EMAIL_DOMAIN=example.com +EMAIL_API=http://127.0.0.1:8080 +``` + +如果 Worker 需要回调本机服务,`WEBHOOK_URL` 应使用可访问的域名地址。 + +## 配置 + +完整模板见 `.env.example`。日常使用通常只需要配置邮箱模式、代理、目标数量和内存预算。 + +| 配置 | 默认值 | 说明 | +|---|---:|---| +| `EMAIL_MODE` | `tempmail` | 邮箱模式,支持 `tempmail` 和 `custom` | +| `EMAIL_DOMAIN` | 空 | `custom` 模式使用的域名 | +| `EMAIL_API` | `http://127.0.0.1:8080` | 本地收信服务地址 | +| `TARGET` | `0` | 成功数量目标,`0` 表示不限 | +| `PHYSICAL_CAP` | `0` | 浏览器并发上限,`0` 表示启动时自动估算 | +| `PHYSICAL_PER_CPU` | `2` | 自动估算时每个 CPU 核心对应的并发参考值 | +| `PHYSICAL_MEM_MB` | `512` | 自动估算时每个浏览器任务的内存预算 | +| `MIN_FREE_MEM_MB` | `500` | 自动估算时保留的内存 | +| `T_SLOT_CAP` | `8` | token 缓冲容量 | +| `Q_SLOT_CAP` | `8` | 验证码缓冲容量 | +| `Q_PENDING_CAP` | `12` | 等待验证码返回的请求上限 | +| `SOLVER_MOUSE_CLICK_RETRIES` | `3` | token 验证框中心点击次数,`0` 表示关闭 | +| `PAGE_BLOCK_STATIC_ASSETS` | `0` | 可选:阻断部分静态资源,降低页面准备成本 | +| `C_HOT_PAGE_POOL` | `0` | 可选:复用消费阶段页面,减少页面重建开销 | + +不确定怎么设置时,先保持默认值。性能压测时优先观察 `PHYSICAL_CAP` 和内存,不建议先改 Worker 数量。 + +## 运行日志 + +直接运行 `bash start.sh` 时,终端只输出任务开始、成功或失败、本次运行平均速度、累计数量和限流等待: + +```text +[→] 开始注册 #38 +[✓] 注册成功 #38 | 运行平均 47/分 | 累计 38 +[⏸] 触发限流 | 60秒后恢复探测 +[▶] 限流解除 | 实际等待 61秒 +``` + +需要调试并发、库存和阶段耗时时,使用: + +```bash +bash start.sh --debug +``` + +它会在上述任务事件之外,每 8 秒输出一次完整的 T/Q、物理并发和阶段耗时面板。已有自动化环境也可继续使用 `REGISTER_LOG_MODE=debug`。 + +常用字段: + +| 字段 | 含义 | +|---|---| +| `T` | 当前可用 token 数量 | +| `Q` | 当前可用验证码数量 | +| `phys` | 空闲浏览器并发许可 | +| `s_phys` / `p_phys` / `c_phys` | S/P/C 获取浏览器许可的平均等待秒数 / 平均持有秒数 | +| `p_stage` | P 阶段平均耗时:建邮箱 / 准备页面 / 发送请求 | +| `c_stage` | C 阶段平均耗时:拿页面 / 验证码校验 / 注册提交 | +| `c_hot` | C 热页池命中 / 未命中次数 | +| `t_solve_avg` | 平均 token 获取时间 | +| `q_sent` / `q_ret` | 已发送 / 已收到的验证码数量 | +| `pair` | 已配对消费次数 | +| `ok` / `fail` | 成功 / 失败数量 | +| `rate` | 当前累计成功速率 | + +简单判断: + +- `T` 长期为 `0` 且 `Q` 有库存,通常是 token 获取较慢。 +- `Q` 长期为 `0` 且 `T` 有库存,通常是邮箱或验证码链路较慢。 +- `phys` 长期为 `0`,说明浏览器并发已经用满。 +- `t_solve_avg` 明显升高,通常表示浏览器压力、网络质量或 token 服务响应变慢。 + +可以用日志分析工具解析已有日志: + +```bash +python3 - <<'PY' +from pathlib import Path +from tools.runtime_log_analyzer import analyze_text +print(analyze_text(Path("run.log").read_text())) +PY +``` + +## 输出文件 + +成功结果写入: + +```text +keys/accounts.txt +keys/grok.txt +``` + +`accounts.txt` 每行格式: + +```text +email:password:sso_token +``` + +`keys/` 目录包含运行结果,默认不会提交到 Git。 + +认证成功后的 OAuth 出货默认在: + +```text +~/Downloads/grok-free-register-auth/authenticated/ +``` + +可用 `keys/async_auth.sh` 把出货整备成 `keys/cpa_ready/` 并同步探活通过的 `access_token` 到 `keys/acc.md`。 + +## 清零流水线(删库跑路) + +只删 `keys/acc.md` / `accounts.txt` **不够**:`acpa_watchdog` 会从 +`~/Downloads/grok-free-register-auth/authenticated/` 把旧号再整备进 +`keys/cpa_ready/`,`sync_acc` 再写回 `acc.md`。状态文件 `_state.tsv` 也会残留幽灵 alive。 + +**全量清零**用项目根的 `reset_pipeline.sh`(默认 dry-run,必须加 `--yes` 才真删): + +```bash +# 建议先停 register / auth-service / async_auth +bash reset_pipeline.sh # dry-run,只列将删内容 +bash reset_pipeline.sh --yes # 真删:注册源 + 认证出货 + cpa_ready + state +bash reset_pipeline.sh --yes --keep-register + # 只清认证/CPA,保留 keys 里新注册的 SSO 源 +bash reset_pipeline.sh --yes --keep-cpa + # 只清认证账本/源快照,保留 cpa_ready(少用) +bash reset_pipeline.sh --yes --wipe-salt + # 连 enrollment ledger salt 一并清(极少需要) +``` + +会清掉的大致范围: + +| 范围 | 路径 | +|------|------| +| 注册源 | `keys/accounts.txt`、`keys/grok.txt`、`keys/auth-sessions.jsonl`、`keys/acc.md` | +| CPA | `keys/cpa_ready/xai-*.json`、`_state.tsv`、`_discarded/` | +| 认证出货 | `$XAI_AUTH_SERVICE_LOCAL_DIR`(默认 `~/Downloads/grok-free-register-auth`)下的 `authenticated/`、`claimed/`、`enrollment-ledger.db*`、`source-snapshot.jsonl` | + +认证目录可用环境变量覆盖: + +```bash +export XAI_AUTH_SERVICE_LOCAL_DIR=~/Downloads/grok-free-register-auth +# 或 +export XAI_ENROLLER_LOCAL_AUTH_DIR=~/Downloads/grok-free-register-auth +``` + +**只清历史归档、不动现网号池**用 `scripts/clean_history.sh`: + +```bash +bash scripts/clean_history.sh # dry-run +bash scripts/clean_history.sh --yes # 清 zip / _discarded / logs / 杂包 +bash scripts/clean_history.sh --yes --deep # 再清 source-snapshot 等更深一层 +``` + +清完后重开: + +```bash +bash start.sh # 或远端注册机继续跑 +bash auth-service.sh # device-flow → authenticated/ +bash keys/async_auth.sh # acpa_watchdog + sync_acc +``` + +## 项目结构 + +```text +grok_register/ 注册核心与 custom 邮箱服务 +xai_enroller/ OAuth 认证服务 +keys/ acpa_watchdog / sync_acc / async_auth(运行时产物 gitignore) +cloudflare/email-worker.js Cloudflare Email Routing Worker 示例 +start.sh 首次配置和运行 +auth-service.sh 认证服务入口 +reset_pipeline.sh 全量清零注册→认证→CPA 运行时数据 +scripts/clean_history.sh 只清历史归档,不动现网号池 +scripts/push_keys_to_auth.sh 本机 SSO 源推远端 auth(需显式 AUTH_SSH_*) +setup.sh 安装依赖 +.env.example 配置模板 +tools/ 日志分析 + GrokSession→CPA/sub2api 前端转换 +tests/ 自动化测试 +docs/architecture.md 并发架构说明 +``` + +## 测试 + +测试依赖与运行依赖分开安装: + +```bash +.venv/bin/pip install -r tests/requirements.txt +``` + +快速检查: + +```bash +python3 -m unittest tests.test_admission_gate tests.test_register_runtime_unittest tests.test_inventory_unittest tests.test_runtime_log_analyzer -v +``` + +完整测试: + +```bash +python3 -m pytest tests -q +``` + +## xAI OAuth Enroller + +`xai_enroller/` 用于把已有 xAI 账号的 SSO 会话转换为 OAuth 凭据,并导入 CPA +凭据库。它独立于注册流程运行:注册机不需要启动,已有账号也不需要重新注册。 + +### 默认同机模式 + +注册和认证在同一个项目目录运行时,无需配置来源: + +```bash +bash auth-service.sh +``` + +认证服务默认读取本项目的 `keys/auth-sessions.jsonl` 和历史 +`keys/accounts.txt`,每 30 秒生成一次经过校验的原子快照。注册仍可同时运行;认证服务 +只读取完整记录,不会读取正在追加的半行。 + +### 分离设备的 SSH 模式 + +注册机和认证服务分开运行时,服务器继续写入注册结果,本地认证服务 +每 30 秒通过一次性 SSH 导出全量 JSONL 快照。快照经逐行校验、`fsync` 后原子替换到 +`~/Downloads/grok-free-register-auth/source-snapshot.jsonl`;同步失败会继续使用上一份有效快照。 +认证使用本机 CloakBrowser Chromium, +成功结果写入 `~/Downloads/grok-free-register-auth/authenticated/`,运行状态、同步快照与 +认证文件分开保存;认证文件格式可以直接供 CPA 使用。 + +先把导出器同步到注册机的 `scripts/` 目录: + +```bash +scp scripts/export_registered_sessions.py user@your-server:/opt/grok-free-register/scripts/ +``` + +然后在本地配置 SSH 连接: + +下面这些变量既可在终端 `export`,也可写入由 `.env.example` 复制出的 `.env`; +`auth-service.sh` 会自动读取该文件。 + +```bash +export XAI_AUTH_SERVICE_SSH_HOST=user@your-server +export XAI_AUTH_SERVICE_SSH_IDENTITY=/path/to/ssh-key.pem # 使用 ssh-agent 时可省略 +export XAI_AUTH_SERVICE_REMOTE_ROOT=/opt/grok-free-register +export XAI_AUTH_SERVICE_SYNC_SEC=30 +``` + +存在 `XAI_AUTH_SERVICE_SSH_HOST` 时会自动选择 SSH。也可显式设置 +`XAI_AUTH_SERVICE_SOURCE=local|ssh`;默认值 `auto` 优先兼容已有 SSH 配置,否则使用同机来源。 + +启动认证服务: + +```bash +bash auth-service.sh +``` + +需要查看队列、重试、节拍和冷却探针时使用: + +```bash +bash auth-service.sh --debug +``` + +默认每次认证至少间隔 10 秒;实测该值比无间隔运行有更高的长期平均成功速率。限流后 +每 60 秒只进行一次恢复探测。需要覆盖时使用: + +```bash +export XAI_AUTH_SERVICE_MIN_INTERVAL_SEC=10 +export XAI_AUTH_SERVICE_RETRY_SEC=60 +``` + +终端只在账号开始、认证结果、限流状态或控制状态变化时输出,并在底部保持 `认证> ` 输入行。`s` 查看状态,`p` 暂停, +`r` 恢复,`c` 取消当前账号,`q` 退出;`Ctrl-C` 同样会退出。在输入行键入 `take 100` +会把最新的 100 个可用凭证登记为已取用,并移动到独立批次目录。认证记录仍保留为 +`imported`,不会因为凭证被取出而重新认证。 + +可用凭证保存在: + +```text +~/Downloads/grok-free-register-auth/authenticated/ +``` + +已取用批次保存在: + +```text +~/Downloads/grok-free-register-auth/claimed// +``` + +库存状态保存在 `enrollment-ledger.db` 的 `credential_inventory` 表中,状态为 +`available`、`claiming` 或 `claimed`。每条记录预留 `note` 字段,默认留空。 + +`auth-service.sh` 首次运行会自动安装项目依赖。正式用户流程只需要这个 Bash 入口;底层 Python 模块保留给开发和测试,不作为另一套使用方式。 + +## 开发文档 + +[docs/architecture.md](docs/architecture.md) 记录并发模型、资源生命周期和必须保持的不变量。 + +## License + +MIT diff --git a/auth-service.sh b/auth-service.sh new file mode 100755 index 0000000..8eaf0ef --- /dev/null +++ b/auth-service.sh @@ -0,0 +1,55 @@ +#!/usr/bin/env bash +set -euo pipefail +cd "$(dirname "$0")" + +. scripts/ensure_runtime.sh +ensure_runtime + +# --------------------------------------------------------------------------- +# 后台拉起 CPA 整备 + acc token 汇总 +# acpa_watchdog.py → authenticated/ → keys/cpa_ready/ (整备 cli-chat-proxy 凭据) +# sync_acc.py → keys/cpa_ready/ → keys/acc.md (只抽 access_token, 一行一个) +# 两者都常驻轮询。不用 exec: exec 会替换 shell, 后台子进程变孤儿且 trap 失效。 +# 脚本退出(正常 / Ctrl-C / SIGTERM)时一并回收, 避免孤儿。 +# --------------------------------------------------------------------------- +WATCHDOG="keys/acpa_watchdog.py" +SYNC_ACC="keys/sync_acc.py" +WATCHDOG_PID="" +SYNC_PID="" + +cleanup() { + for pid in "${SYNC_PID:-}" "${WATCHDOG_PID:-}"; do + if [ -n "$pid" ] && kill -0 "$pid" 2>/dev/null; then + kill "$pid" 2>/dev/null || true + wait "$pid" 2>/dev/null || true + fi + done +} +trap cleanup EXIT INT TERM + +mkdir -p logs + +if [ -f "$WATCHDOG" ]; then + .venv/bin/python "$WATCHDOG" >logs/watchdog.log 2>&1 & + WATCHDOG_PID=$! + echo "[auth-service] watchdog 已拉起 (pid=$WATCHDOG_PID, 日志 logs/watchdog.log)" +else + echo "[auth-service] 未找到 $WATCHDOG, 跳过自动整备" >&2 +fi + +if [ -f "$SYNC_ACC" ]; then + .venv/bin/python "$SYNC_ACC" >logs/sync_acc.log 2>&1 & + SYNC_PID=$! + echo "[auth-service] sync_acc 已拉起 (pid=$SYNC_PID, 日志 logs/sync_acc.log → keys/acc.md)" +else + echo "[auth-service] 未找到 $SYNC_ACC, 跳过 acc.md 汇总" >&2 +fi + +# --------------------------------------------------------------------------- +# 前台跑认证服务 (注册出货)。退出码透传; trap cleanup 在 exit 时兜底回收后台进程。 +# --------------------------------------------------------------------------- +set +e +.venv/bin/python -m xai_enroller.service "$@" +RC=$? +set -e +exit "$RC" diff --git a/cloudflare/email-worker.js b/cloudflare/email-worker.js new file mode 100644 index 0000000..52bb4d5 --- /dev/null +++ b/cloudflare/email-worker.js @@ -0,0 +1,61 @@ +// ============================================================ +// Cloudflare Email Worker — 把收到的邮件 POST 到你的 webhook +// 用于 grok-free-register 的「自建邮箱模式」(EMAIL_MODE=custom) +// +// 链路: 发件方 → Cloudflare Email Routing → 本 Worker → 你的 webhook(email_server.py) +// +// ⚠️ 重要(踩过的坑):env.WEBHOOK_URL 必须是【域名】,不能用【裸 IP】! +// Cloudflare Workers 的 fetch 不允许直接请求裸 IP,会返回 +// "error code: 1003 (Direct IP access not allowed)" 并被静默吞掉 +// —— Worker 看起来成功、你的服务器却收不到任何请求。 +// 做法:给服务器 IP 加一条 DNS A 记录(如 hook.example.com,灰云/DNS-only), +// WEBHOOK_URL 填 http://hook.example.com:8080/webhook。 +// +// 部署: +// npm create cloudflare@latest cf-mail-webhook +// cd cf-mail-webhook && npm i postal-mime +// # 用本文件替换 src/index.js +// npx wrangler secret put WEBHOOK_URL # 输入 http://hook.example.com:8080/webhook +// npx wrangler secret put WEBHOOK_TOKEN # 可选,鉴权用 +// npx wrangler deploy +// 然后在 Cloudflare 后台:Email → Email Routing → Routing rules → +// Catch-all → 动作选「Send to a Worker」→ 选本 Worker。 +// (子域名收信需在 Email Routing 单独启用该子域并配它自己的 catch-all。) +// ============================================================ +import PostalMime from "postal-mime"; + +function trim(s, max = 20000) { + if (!s) return ""; + return s.length > max ? s.slice(0, max) + "\n...[truncated]" : s; +} + +export default { + async email(message, env, ctx) { + if (!env.WEBHOOK_URL) { + console.error("WEBHOOK_URL 未配置 (npx wrangler secret put WEBHOOK_URL)"); + return; + } + const parsed = await PostalMime.parse(message.raw); + const payload = { + from: message.from, + to: message.to, + subject: parsed.subject || message.headers.get("subject") || "", + text: trim(parsed.text || ""), + html: trim(parsed.html || ""), + }; + + const headers = { "content-type": "application/json" }; + if (env.WEBHOOK_TOKEN) headers["x-webhook-token"] = env.WEBHOOK_TOKEN; + + const res = await fetch(env.WEBHOOK_URL, { + method: "POST", + headers, + body: JSON.stringify(payload), + }); + // 日志只保留状态码,避免把 webhook 地址或上游响应写入 Worker 日志。 + console.log(`webhook delivery -> ${res.status}`); + if (!res.ok) { + console.error(`webhook delivery failed: ${res.status}`); + } + }, +}; diff --git a/docs/architecture.md b/docs/architecture.md new file mode 100644 index 0000000..a729cbd --- /dev/null +++ b/docs/architecture.md @@ -0,0 +1,180 @@ +# CSP 架构说明 + +本文档记录当前运行时架构。README 面向使用者;本文面向维护者,重点是并发边界、资源生命周期和测试不变量。 + +## 目标 + +运行时采用 CSP 风格的异步流水线: + +- `S_Worker` 生产 `T`。 +- `P_Worker` 发起外部请求,等待并生产 `Q`。 +- `C_Worker` 原子获取一组 `T + Q` 并执行最终消费。 + +架构目标是资源所有权闭合、背压边界清晰、取消语义可测试。当前设计不包含中心调度器、动态角色选择或运行时动态并发分配。 + +## 核心组件 + +`Physical_Sem` 限制本地浏览器重操作并发。`P_Worker` 发出请求后,在等待 `Q` 返回期间不得持有该许可。 + +`T_Slot_Sem` 限制已入库 `T` 的容量。`T` 生成成功后才允许获取 slot。 + +`Q_Slot_Sem` 限制已返回并入库 `Q` 的容量。`Q` 真正返回前不得获取 slot。 + +`Q_Pending_Sem` 限制已发出但尚未终态的外部 `Q` 请求数量。 + +`Inventory` 是唯一库存门面。worker 不直接访问底层队列,只能调用: + +```text +put_t(env) +put_q(env) +claim_pair() +``` + +`ResourceEnvelope` 把资源实体和库存 slot 绑定在一起。slot 只能释放一次。 + +`PairLease` 是 `claim_pair()` 返回的异步上下文管理器。pair 一旦 claim 成功,两个 envelope 的所有权转移给 lease,直到 consumer 退出上下文。 + +`AdmissionGate` 是局部生产准入门控,只根据库存深度和静态水位决定是否允许继续生产。它不选择 worker 角色,不搬运资源,不调整并发容量。 + +## Worker 流程 + +### S_Worker + +1. 等待 `AdmissionGate` 允许生产 `T`。 +2. 获取 `Physical_Sem`。 +3. 生产 `T`。 +4. 释放 `Physical_Sem`。 +5. 调用 `ResourceEnvelope.create_with_slot(...)`,创建 envelope 并获取 `T_Slot_Sem`。 +6. 调用 `Inventory.put_t(...)`,把所有权转移给 `Inventory`。 +7. 所有权转移前如果异常或取消,释放已获取的 slot。 + +slot 获取和 envelope 创建必须绑定,避免取消落在“已获取 slot、尚未创建 envelope”的窗口。 + +### P_Worker + +1. 等待 `AdmissionGate` 允许生产 `Q`。 +2. 获取 `Q_Pending_Sem`。 +3. 获取 `Physical_Sem`。 +4. 创建请求并发送。 +5. 释放 `Physical_Sem`。 +6. 在不持有 `Physical_Sem` 的情况下等待 `Q` 返回。 +7. `Q` 返回后调用 `ResourceEnvelope.create_with_slot(...)`,创建 envelope 并获取 `Q_Slot_Sem`。 +8. 调用 `Inventory.put_q(...)`,把所有权转移给 `Inventory`。 +9. 请求进入终态后释放 `Q_Pending_Sem`。 + +`Q_Pending_Sem` 表达外部在途上限;`Q_Slot_Sem` 只表达已返回库存容量。两者不能合并。 + +### C_Worker + +1. 进入 `async with inventory.claim_pair() as pair`。 +2. 获取 `Physical_Sem`。 +3. 消费 pair。 +4. 释放 `Physical_Sem`。 +5. 退出 `PairLease`,释放两个库存 slot。 + +`C_Worker` 不允许先取单边资源再等待另一边。对 consumer 来说,pair claim 必须是原子的。 + +## Inventory 语义 + +第一版 `Inventory` 使用一把 lock 和一个 condition。lock 保护等待、复查、过期清理和弹出操作。 + +必须保持以下语义: + +- 等待中的 consumer 不移除资源。 +- claim 成功时同时移除一个有效 `T` 和一个有效 `Q`。 +- 等待 pair 时被取消,不影响库存。 +- claim 成功后被取消,由 `PairLease` 核销两个 envelope。 +- worker 永远不能直接操作底层 `T` / `Q` 队列。 + +只要锁内逻辑保持很小,除 lazy expiry cleanup 外基本是 O(1),单锁在当前版本可以接受。只有 profile 证明锁竞争成为真实瓶颈时,才考虑拆锁或分片。 + +## 过期模型 + +`ResourceEnvelope` 可以携带 `created_at` 和 `expires_at`。`Inventory` 在配对前可以丢弃已过期资源。 + +当前采用 lazy cleanup: + +- `put_t`、`put_q`、`claim_pair` 被触发时顺手清理。 +- 清理会释放它看到的过期 envelope 对应 slot。 +- 系统完全静默时不会主动扫库。 +- 单边长期故障和静默停摆由监控暴露,不靠后台 sweeper 修复。 + +当前不做 worker 级回队。消费失败后,已 claim 的 `T` 和 `Q` 都由 `PairLease` 核销。 + +## 容量策略 + +容量边界由 Semaphore 表达。启动期容量优先级: + +```text +显式 PHYSICAL_CAP > CAPACITY_PROFILE > CPU/内存自动派生 +``` + +`CAPACITY_PROFILE` 只在启动时读取,是静态 profile,不是运行时调度器。 + +默认 worker 数量由容量派生: + +```text +S_WORKERS = Physical_Sem + 2 +P_WORKERS = Q_Pending_Sem + 2 +C_WORKERS = Physical_Sem + 2 +``` + +worker 数量不是主要调参入口。主要并发边界是各类容量许可,而不是 coroutine 循环数量。 + +## 当前不做 + +当前版本明确不包含: + +- 中心调度器; +- 运行时角色选择; +- 动态打分; +- 动态并发控制; +- worker 级回队; +- 高价值资源抢救策略; +- 后台过期清扫; +- 自动切换高风险浏览器模式。 + +这些方向可以单独实验,但不能混入基础所有权模型。 + +## 必须保持的不变量 + +实现和测试必须维持以下不变量: + +- 每个已获取的库存 slot 最终释放一次且只释放一次。 +- 每个已准入的 pending 请求最终释放一次 `Q_Pending_Sem`。 +- `P_Worker` 等待 `Q` 返回时不持有 `Physical_Sem`。 +- `Q_Slot_Sem` 只在 `Q` 返回后获取。 +- `C_Worker` 只能通过 `Inventory.claim_pair()` 获取 `T` 和 `Q`。 +- `claim_pair()` 要么返回一个受 `PairLease` 保护的完整 pair,要么不返回资源。 +- 等待 pair 时取消,不移除库存资源。 +- claim pair 后取消,两个 envelope 由 `PairLease` 释放。 +- 触发清理时,过期资源不能被配对。 +- 监控只读,不修改 Semaphore、队列或 worker 状态。 + +## 测试 + +```bash +.venv/bin/pip install -r tests/requirements.txt +``` + +快速检查: + +```bash +python3 -m unittest tests.test_admission_gate tests.test_register_runtime_unittest tests.test_inventory_unittest tests.test_runtime_log_analyzer -v +``` + +完整测试: + +```bash +python3 -m pytest tests -q +``` + +重点测试文件: + +- `tests.test_inventory_unittest`:库存、过期和 lease 行为。 +- `tests.test_register_runtime_unittest`:worker 运行语义和监控行。 +- `tests.test_admission_gate`:局部门控水位。 +- `tests.test_runtime_log_analyzer`:日志解析兼容性。 +- `tests/test_cancel.py`:取消边界。 +- `tests/test_property.py`:随机化不变量检查。 +- `tests/test_stress.py`:更高并发 fake-service 压测。 diff --git a/docs/guides/auth-service.md b/docs/guides/auth-service.md new file mode 100644 index 0000000..2349977 --- /dev/null +++ b/docs/guides/auth-service.md @@ -0,0 +1,67 @@ +# 本地认证服务 + +认证服务把已有的 SSO 会话转换为 CPA 可直接读取的 OAuth 凭据。注册与认证可以在同一台机器,也可以分开运行。 + +## 默认同机运行 + +同一个项目目录已经或正在运行注册服务时,直接启动认证: + +```bash +bash auth-service.sh +``` + +未配置 SSH 主机时,认证服务自动读取本项目 `keys/` 中的完整会话与历史账号。注册可以继续追加,认证服务只安装经过校验的完整快照。 + +## 配置远端同步 + +先把无密码导出器放到服务器项目目录: + +```bash +scp scripts/export_registered_sessions.py user@server.example:/opt/grok-free-register/scripts/ +``` + +在本地终端设置连接信息: + +可以直接 `export`,也可以把 `.env.example` 复制为 `.env` 后填写;认证入口会自动读取 `.env`。 + +```bash +export XAI_AUTH_SERVICE_SSH_HOST=user@server.example +export XAI_AUTH_SERVICE_SSH_IDENTITY=/path/to/key.pem +export XAI_AUTH_SERVICE_REMOTE_ROOT=/opt/grok-free-register +``` + +使用 `ssh-agent` 时可省略 `XAI_AUTH_SERVICE_SSH_IDENTITY`。 + +设置了 `XAI_AUTH_SERVICE_SSH_HOST` 后,默认的 `auto` 模式会选择 SSH。需要明确覆盖时使用: + +```bash +export XAI_AUTH_SERVICE_SOURCE=local # 强制读取同机注册结果 +export XAI_AUTH_SERVICE_SOURCE=ssh # 强制使用 SSH,必须配置主机 +``` + +## 运行 + +```bash +bash auth-service.sh +``` + +首次运行会自动安装项目依赖。该命令在当前终端持续运行并直接接受控制命令;输入 `q` 或按 `Ctrl-C` 停止,再次执行同一命令即可重启。不需要额外的会话管理工具。 + +普通模式只在来源连接、发现新账号、任务开始、认证结果、限流和控制状态变化时输出。查看队列、重试、节拍和冷却探针时使用: + +```bash +bash auth-service.sh --debug +``` + +运行中终端底部会保持 `认证> ` 输入行。日志更新不会清掉尚未提交的内容;直接输入命令并回车: + +```text +s 查看状态 +take N 取用 N 个凭据 +p 暂停 +r 恢复 +c 取消当前任务 +q 安全退出 +``` + +快照默认每 30 秒更新一次;内容无变化时终端保持安静。有效快照和已生成凭据会在重启后继续使用。 diff --git a/docs/guides/credential-inventory.md b/docs/guides/credential-inventory.md new file mode 100644 index 0000000..4db1227 --- /dev/null +++ b/docs/guides/credential-inventory.md @@ -0,0 +1,19 @@ +# 凭据库存与取用 + +认证成功的凭据保存在本地认证目录,并在 SQLite 库存中维护三种状态: + +- `available`:已经认证、尚未取用; +- `claiming`:正在移动到取用批次; +- `claimed`:已经取用。 + +运行认证服务后输入: + +```text +take 100 +``` + +服务会选择最新的 100 个可用凭据,移动到 `claimed//`,再把对应库存标记为 `claimed`。认证 ledger 中的 `imported` 记录仍然保留,所以这些账号不会被重新认证。 + +每条库存记录预留 `note` 字段,默认为空。取用操作不要求填写用途;以后需要备注时可直接更新该字段。 + +若文件移动中断,服务下次启动会恢复 `claiming` 状态。库存不足或凭据文件缺失时,操作失败但认证服务继续运行。 diff --git a/docs/guides/registration.md b/docs/guides/registration.md new file mode 100644 index 0000000..d2b0b25 --- /dev/null +++ b/docs/guides/registration.md @@ -0,0 +1,54 @@ +# 注册教程 + +## 开始运行 + +```bash +git clone /grok-free-register.git +cd grok-free-register +bash start.sh +``` + +首次运行会安装 Python、CloakBrowser Chromium 及其系统依赖,然后引导选择邮箱模式。以后再次执行 `bash start.sh` 会直接使用已有配置。 + +普通模式只显示服务启动、任务开始、注册成功或失败、本次运行平均速率、累计数量和限流状态。查看完整并发、库存和阶段耗时时使用: + +```bash +bash start.sh --debug +``` + +常用参数: + +```bash +bash start.sh --target 100 +bash start.sh --max-mem 6G +bash start.sh --reconfig +``` + +未设置 `--target` 时服务持续运行,按 `Ctrl-C` 安全停止。 +再次执行 `bash start.sh` 即可重启。程序直接使用当前终端,不需要额外的会话管理工具。 + +## 配置邮箱 + +临时邮箱无需额外配置: + +```env +EMAIL_MODE=tempmail +``` + +自建邮箱需要可接收邮件的域名和本项目的收信服务: + +```env +EMAIL_MODE=custom +EMAIL_DOMAIN=example.com +EMAIL_API=http://127.0.0.1:8080 +``` + +自建模式还需运行: + +```bash +bash start.sh --email-service +``` + +性能参数默认会根据 CPU 和可用内存估算。除非正在压测,否则保持 `.env.example` 中的默认值即可。 + +成功结果写入 `keys/accounts.txt`、`keys/grok.txt` 和 `keys/auth-sessions.jsonl`;这些文件默认不提交到 Git。 diff --git a/docs/guides/runtime-troubleshooting.md b/docs/guides/runtime-troubleshooting.md new file mode 100644 index 0000000..6689027 --- /dev/null +++ b/docs/guides/runtime-troubleshooting.md @@ -0,0 +1,38 @@ +# 运行状态与排障 + +## 普通模式 + +普通模式的每一行都对应一次状态变化: + +- `[→]`:新任务开始; +- `[✓]`:任务成功或来源连接完成; +- `[✗]`:当前任务失败或被跳过; +- `[⏸]`:进入限流冷却或服务暂停; +- `[▶]`:限流解除或服务恢复; +- `[!]`:来源、配置或流水线出现需要关注的问题。 + +认证端输入 `s` 可查看运行状态、待处理数量、当前阶段、本次运行平均速率、累计成功、可用和已取用凭据,以及是否处于限流。 + +## Debug 模式 + +注册端: + +```bash +bash start.sh --debug +``` + +认证端: + +```bash +bash auth-service.sh --debug +``` + +注册 Debug 面板包含 T/Q 库存、物理并发、S/P/C 阶段耗时和 token 求解时间。认证 Debug 状态包含 source/prepared/completion 队列、重试、授权节拍、冷却、单次探针和近五分钟滚动速率。 + +## 常见状态 + +限流后不会持续重试。认证端默认等待 60 秒,再只放行一个恢复探针;探针仍被限流时重新等待。注册端同样通过全局冷却闸门阻止并发任务漏过等待周期。 + +远端来源暂时断开时,本地认证服务继续使用上一份完整有效快照。恢复连接后会自动同步,不需要重启。 + +配置错误会指出缺少或非法的配置名,不输出 traceback。按提示检查 [注册配置](registration.md#配置邮箱) 或 [认证同步配置](auth-service.md#配置远端同步)。 diff --git a/grok_register/__init__.py b/grok_register/__init__.py new file mode 100644 index 0000000..91e7867 --- /dev/null +++ b/grok_register/__init__.py @@ -0,0 +1 @@ +"""Registration service package.""" diff --git a/grok_register/clearance.py b/grok_register/clearance.py new file mode 100644 index 0000000..19feb7e --- /dev/null +++ b/grok_register/clearance.py @@ -0,0 +1,395 @@ +"""Cloudflare clearance prewarm via external FlareSolverr. + +This module does NOT run FlareSolverr. It only calls an already-deployed +FlareSolverr instance (default http://127.0.0.1:8191). + +Proxy roles (do not mix): + - REGISTER_PROXY / HTTP_PROXY / HTTPS_PROXY + → grok-free-register Playwright + httpx egress (this process) + - CLEARANCE_PROXY + → optional proxy *inside* FlareSolverr's browser (must be reachable + from the FS *container*, e.g. http://privoxy:8118) + - FLARESOLVERR_URL + → host-side URL to reach FS (http://127.0.0.1:8191) + +Clearance cookies are only valid on the same egress path they were minted on. +If REGISTER_PROXY points at host Privoxy→WARP, set CLEARANCE_PROXY to the +docker-internal privoxy URL so FS solves on the same WARP exit. +""" + +from __future__ import annotations + +import json +import os +import threading +import time +import urllib.error +import urllib.request +from typing import Any +from urllib.parse import urlparse + +# CF-fronted x.ai family roots (no path suffix; path pages share host clearance) +DEFAULT_CLEARANCE_URLS = ( + "https://accounts.x.ai", + "https://x.ai", + "https://status.x.ai", + "https://console.x.ai", + "https://auth.x.ai", +) + +_lock = threading.RLock() +_cache: dict[str, Any] = { + "cookies": [], # Playwright cookie dicts + "user_agent": "", + "fetched_at": 0.0, + "hosts": [], + "last_error": "", +} + + +def _env_bool(name: str, default: bool = False) -> bool: + raw = str(os.environ.get(name, "")).strip().lower() + if not raw: + return default + if raw in {"1", "true", "yes", "on"}: + return True + if raw in {"0", "false", "no", "off"}: + return False + return default + + +def _env_int(name: str, default: int) -> int: + try: + return int(str(os.environ.get(name, "")).strip() or default) + except (TypeError, ValueError): + return default + + +def clearance_enabled() -> bool: + return _env_bool("CLEARANCE_ENABLED", False) + + +def flaresolverr_url() -> str: + return (os.environ.get("FLARESOLVERR_URL") or "http://127.0.0.1:8191").strip().rstrip("/") + + +def clearance_timeout_sec() -> int: + return max(10, _env_int("CLEARANCE_TIMEOUT_SEC", 60)) + + +def clearance_refresh_sec() -> int: + return max(60, _env_int("CLEARANCE_REFRESH_SEC", 3000)) + + +def clearance_urls() -> list[str]: + raw = (os.environ.get("CLEARANCE_URLS") or "").strip() + if not raw: + return list(DEFAULT_CLEARANCE_URLS) + urls = [] + for part in raw.split(","): + item = part.strip() + if not item: + continue + if "://" not in item: + item = "https://" + item + # strip path/query — host-level prewarm only + parsed = urlparse(item) + if parsed.scheme and parsed.netloc: + urls.append(f"{parsed.scheme}://{parsed.netloc}") + else: + urls.append(item.rstrip("/")) + # de-dupe preserve order + seen = set() + out = [] + for u in urls: + if u not in seen: + seen.add(u) + out.append(u) + return out or list(DEFAULT_CLEARANCE_URLS) + + +def register_proxy_url() -> str: + """Egress proxy for *this* process (Playwright / httpx).""" + for key in ("REGISTER_PROXY", "HTTPS_PROXY", "HTTP_PROXY", "ALL_PROXY"): + value = (os.environ.get(key) or "").strip() + if value: + return value + return "" + + +def clearance_proxy_url() -> str: + """Proxy for FlareSolverr browser (container-reachable). Empty = FS direct egress.""" + explicit = (os.environ.get("CLEARANCE_PROXY") or "").strip() + if explicit: + return explicit + # Do NOT auto-forward host REGISTER_PROXY into FS — 127.0.0.1 inside + # the container is not the host. Operator must set CLEARANCE_PROXY when + # minting on WARP (e.g. http://privoxy:8118). + return "" + + +def playwright_proxy_settings() -> dict[str, str] | None: + """Playwright launch/context proxy dict, or None for direct.""" + url = register_proxy_url() + if not url: + return None + return {"server": url} + + +def httpx_proxy_mounts() -> str | None: + """Single proxy URL for httpx.AsyncClient(proxy=...), or None.""" + return register_proxy_url() or None + + +def _host_from_url(url: str) -> str: + try: + return (urlparse(url).hostname or "").lower() + except Exception: + return "" + + +def _fs_post(payload: dict[str, Any], timeout: float) -> dict[str, Any]: + endpoint = f"{flaresolverr_url()}/v1" + body = json.dumps(payload).encode("utf-8") + req = urllib.request.Request( + endpoint, + data=body, + headers={"Content-Type": "application/json"}, + method="POST", + ) + with urllib.request.urlopen(req, timeout=timeout) as response: + raw = response.read() + return json.loads(raw.decode("utf-8")) + + +def _cookie_to_playwright(item: dict[str, Any], fallback_host: str) -> dict[str, Any] | None: + name = str(item.get("name") or "").strip() + if not name: + return None + value = str(item.get("value") or "") + domain = str(item.get("domain") or "").strip() + if not domain: + domain = fallback_host + if domain and not domain.startswith(".") and domain.count(".") >= 1: + # keep host-only cookies host-only; CF often uses .x.ai + pass + path = str(item.get("path") or "/") or "/" + cookie: dict[str, Any] = { + "name": name, + "value": value, + "domain": domain, + "path": path, + } + if "httpOnly" in item: + cookie["httpOnly"] = bool(item.get("httpOnly")) + if "secure" in item: + cookie["secure"] = bool(item.get("secure")) + else: + cookie["secure"] = True + same_site = item.get("sameSite") or item.get("same_site") + if same_site: + # Playwright: Strict | Lax | None + normalized = str(same_site).capitalize() + if normalized.lower() == "none": + normalized = "None" + if normalized in {"Strict", "Lax", "None"}: + cookie["sameSite"] = normalized + expires = item.get("expires") + if expires is not None: + try: + exp = float(expires) + # FS may return ms or session -1 + if exp > 0: + if exp > 1e12: + exp = exp / 1000.0 + cookie["expires"] = exp + except (TypeError, ValueError): + pass + return cookie + + +def _merge_cookies(existing: list[dict[str, Any]], incoming: list[dict[str, Any]]) -> list[dict[str, Any]]: + index: dict[tuple[str, str, str], int] = {} + merged: list[dict[str, Any]] = [] + for cookie in existing: + key = ( + str(cookie.get("name") or ""), + str(cookie.get("domain") or ""), + str(cookie.get("path") or "/"), + ) + index[key] = len(merged) + merged.append(cookie) + for cookie in incoming: + key = ( + str(cookie.get("name") or ""), + str(cookie.get("domain") or ""), + str(cookie.get("path") or "/"), + ) + if key in index: + merged[index[key]] = cookie + else: + index[key] = len(merged) + merged.append(cookie) + return merged + + +def prewarm_clearance(*, force: bool = False) -> dict[str, Any]: + """Hit FlareSolverr for each CLEARANCE_URLS host; cache Playwright cookies. + + Safe to call when disabled — returns status without network. + """ + if not clearance_enabled(): + return { + "enabled": False, + "ok": True, + "skipped": True, + "message": "CLEARANCE_ENABLED=0", + "cookies": 0, + } + + with _lock: + age = time.time() - float(_cache.get("fetched_at") or 0) + if ( + not force + and _cache.get("cookies") + and age < clearance_refresh_sec() + ): + return { + "enabled": True, + "ok": True, + "cached": True, + "age_sec": round(age, 1), + "cookies": len(_cache["cookies"]), + "hosts": list(_cache.get("hosts") or []), + "user_agent": _cache.get("user_agent") or "", + } + + urls = clearance_urls() + timeout = float(clearance_timeout_sec()) + fs_proxy = clearance_proxy_url() + max_timeout_ms = int(timeout * 1000) + merged: list[dict[str, Any]] = [] + user_agent = "" + hosts_ok: list[str] = [] + errors: list[str] = [] + per_host: list[dict[str, Any]] = [] + + for url in urls: + host = _host_from_url(url) + payload: dict[str, Any] = { + "cmd": "request.get", + "url": url, + "maxTimeout": max_timeout_ms, + } + if fs_proxy: + payload["proxy"] = {"url": fs_proxy} + t0 = time.time() + try: + data = _fs_post(payload, timeout=timeout + 15) + elapsed = round(time.time() - t0, 2) + if str(data.get("status") or "").lower() != "ok": + errors.append(f"{host}: status={data.get('status')} msg={data.get('message')}") + per_host.append({"host": host, "ok": False, "elapsed": elapsed, "error": data.get("message")}) + continue + solution = data.get("solution") if isinstance(data.get("solution"), dict) else {} + raw_cookies = solution.get("cookies") or [] + pw_cookies = [] + if isinstance(raw_cookies, list): + for item in raw_cookies: + if not isinstance(item, dict): + continue + converted = _cookie_to_playwright(item, host) + if converted: + pw_cookies.append(converted) + merged = _merge_cookies(merged, pw_cookies) + ua = str(solution.get("userAgent") or "").strip() + if ua: + user_agent = ua + has_cf = any(c.get("name") == "cf_clearance" for c in pw_cookies) + hosts_ok.append(host) + per_host.append( + { + "host": host, + "ok": True, + "elapsed": elapsed, + "http": solution.get("status"), + "cookies": len(pw_cookies), + "cf_clearance": has_cf, + } + ) + except Exception as exc: # noqa: BLE001 — surface to operator log + elapsed = round(time.time() - t0, 2) + errors.append(f"{host}: {exc}") + per_host.append({"host": host, "ok": False, "elapsed": elapsed, "error": str(exc)}) + + with _lock: + if merged: + _cache["cookies"] = merged + _cache["user_agent"] = user_agent + _cache["fetched_at"] = time.time() + _cache["hosts"] = hosts_ok + _cache["last_error"] = "; ".join(errors) + elif errors: + _cache["last_error"] = "; ".join(errors) + + return { + "enabled": True, + "ok": bool(merged) or not errors, + "cached": False, + "cookies": len(merged), + "hosts": hosts_ok, + "user_agent": user_agent, + "errors": errors, + "detail": per_host, + "flaresolverr": flaresolverr_url(), + "clearance_proxy": fs_proxy or "(direct)", + "register_proxy": register_proxy_url() or "(direct)", + } + + +def cached_playwright_cookies() -> list[dict[str, Any]]: + with _lock: + return [dict(c) for c in (_cache.get("cookies") or [])] + + +def cached_user_agent() -> str: + with _lock: + return str(_cache.get("user_agent") or "") + + +async def apply_clearance_to_context(context) -> int: + """Inject cached CF cookies into a Playwright BrowserContext. Returns count.""" + cookies = cached_playwright_cookies() + if not cookies or context is None: + return 0 + try: + await context.add_cookies(cookies) + return len(cookies) + except Exception: + return 0 + + +def format_prewarm_log(result: dict[str, Any]) -> str: + if result.get("skipped"): + return "[clearance] 未启用 (CLEARANCE_ENABLED=0)" + if result.get("cached"): + return ( + f"[clearance] 复用缓存 cookies={result.get('cookies')} " + f"age={result.get('age_sec')}s hosts={','.join(result.get('hosts') or [])}" + ) + detail = result.get("detail") or [] + parts = [] + for row in detail: + host = row.get("host") + if row.get("ok"): + cf = "cf+" if row.get("cf_clearance") else "cf-" + parts.append(f"{host}:{row.get('elapsed')}s/{cf}") + else: + parts.append(f"{host}:FAIL") + err = result.get("errors") or [] + suffix = f" errors={len(err)}" if err else "" + return ( + f"[clearance] 预热完成 cookies={result.get('cookies')} " + f"fs={result.get('flaresolverr')} reg_proxy={result.get('register_proxy')} " + f"fs_proxy={result.get('clearance_proxy')} | " + " ".join(parts) + suffix + ) diff --git a/grok_register/core/__init__.py b/grok_register/core/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/grok_register/core/admission.py b/grok_register/core/admission.py new file mode 100644 index 0000000..bc9f8fc --- /dev/null +++ b/grok_register/core/admission.py @@ -0,0 +1,129 @@ +"""Local admission gates for CSP workers. + +AdmissionGate is not a scheduler: it never chooses a role and never moves +resources. It only lets producers reserve local production intent against +inventory watermarks, so fixed S/P workers can self-throttle without a central +dispatcher. +""" +import asyncio + + +class _TProductionLease: + __slots__ = ("_gate", "_released") + + def __init__(self, gate): + self._gate = gate + self._released = False + + async def release(self): + if self._released: + return + self._released = True + async with self._gate._cond: + self._gate.t_in_progress -= 1 + self._gate._cond.notify_all() + + +class _QBatchLease: + __slots__ = ("_gate", "count", "_remaining") + + def __init__(self, gate, count): + self._gate = gate + self.count = count + self._remaining = count + + async def release_one(self): + if self._remaining <= 0: + return + async with self._gate._cond: + if self._remaining <= 0: + return + self._remaining -= 1 + self._gate.q_inflight -= 1 + self._gate._cond.notify_all() + + async def release_all(self): + if self._remaining <= 0: + return + async with self._gate._cond: + if self._remaining <= 0: + return + self._gate.q_inflight -= self._remaining + self._remaining = 0 + self._gate._cond.notify_all() + + +class AdmissionGate: + """Watermark-based producer admission for T and Q. + + T production uses high/low hysteresis because local T generation can be + heavy and should not keep filling an already sufficient inventory. + + Q production reserves a batch by current demand: + q_high - q_depth - q_inflight + capped by the worker's max batch. Each reservation is released when the + corresponding pending Q requests have reached terminal state. + """ + + def __init__(self, inventory, *, t_low, t_high, q_low, q_high): + if t_low < 0 or q_low < 0: + raise ValueError("low watermarks must be non-negative") + if t_high <= 0 or q_high <= 0: + raise ValueError("high watermarks must be positive") + if t_low > t_high: + raise ValueError("t_low must be <= t_high") + if q_low > q_high: + raise ValueError("q_low must be <= q_high") + + self.inventory = inventory + self.t_low = t_low + self.t_high = t_high + self.q_low = q_low + self.q_high = q_high + self.t_in_progress = 0 + self.q_inflight = 0 + self._t_paused = False + self._q_paused = False + self._cond = asyncio.Condition() + + async def acquire_t_production(self): + """Reserve one T production attempt.""" + async with self._cond: + while True: + total = self.inventory.t_depth + self.t_in_progress + if total >= self.t_high: + self._t_paused = True + if self._t_paused and total <= self.t_low: + self._t_paused = False + if not self._t_paused and total < self.t_high: + self.t_in_progress += 1 + self._cond.notify_all() + return _TProductionLease(self) + await self._cond.wait() + + async def acquire_q_batch(self, *, max_batch): + """Reserve up to max_batch pending Q requests based on current demand.""" + if max_batch <= 0: + raise ValueError("max_batch must be positive") + + async with self._cond: + while True: + total = self.inventory.q_depth + self.q_inflight + if total >= self.q_high: + self._q_paused = True + if self._q_paused and total <= self.q_low: + self._q_paused = False + demand = self.q_high - total + if not self._q_paused and demand > 0: + count = min(max_batch, demand) + self.q_inflight += count + if self.inventory.q_depth + self.q_inflight >= self.q_high: + self._q_paused = True + self._cond.notify_all() + return _QBatchLease(self, count) + await self._cond.wait() + + async def notify_changed(self): + """Wake waiters after inventory depth changes.""" + async with self._cond: + self._cond.notify_all() diff --git a/grok_register/core/envelope.py b/grok_register/core/envelope.py new file mode 100644 index 0000000..e3bb587 --- /dev/null +++ b/grok_register/core/envelope.py @@ -0,0 +1,61 @@ +""" +ResourceEnvelope — T/Q 资源实体与库存 slot 的生命周期绑定 + +slot 获取与 envelope 创建必须通过 create_with_slot() 完成。 +release_slot_once() 幂等:重复调用只释放一次。 +""" +import time +import asyncio + + +class ResourceEnvelope: + """绑定一个 T/Q 资源值与其对应的库存 slot。""" + + __slots__ = ('kind', 'value', 'slot_sem', '_released', 'created_at', 'expires_at', 'meta') + + def __init__(self, kind, value, slot_sem, *, created_at, expires_at=None, meta=None): + self.kind = kind # 'T' or 'Q' + self.value = value # T: turnstile token str; Q: {'email','password','code'} + self.slot_sem = slot_sem # T_Slot_Sem or Q_Slot_Sem + self._released = False + self.created_at = created_at + self.expires_at = expires_at + self.meta = meta or {} + + @classmethod + async def create_with_slot(cls, kind, value, slot_sem, *, expires_at=None, meta=None): + """先获取库存 slot,再构造 envelope。slot 获取失败则抛异常,不创建 envelope。""" + await slot_sem.acquire() + try: + return cls( + kind, value, slot_sem, + created_at=time.time(), + expires_at=expires_at, + meta=meta, + ) + except BaseException: + slot_sem.release() + raise + + def release_slot_once(self, reason=''): + """幂等释放 slot。重复调用安全。""" + if not self._released: + self._released = True + self.slot_sem.release() + + def discard(self, reason=''): + """资源未被消费,释放库存占用。""" + self.release_slot_once(reason=f'discard:{reason}') + + def is_expired(self, now=None): + """检查资源是否过期。""" + if self.expires_at is None: + return False + return (now or time.time()) > self.expires_at + + @property + def released(self): + return self._released + + def __repr__(self): + return f'Envelope({self.kind}, released={self._released})' diff --git a/grok_register/core/inventory.py b/grok_register/core/inventory.py new file mode 100644 index 0000000..12068a4 --- /dev/null +++ b/grok_register/core/inventory.py @@ -0,0 +1,169 @@ +""" +Inventory — 唯一库存门面 + PairLease + +put_t / put_q: 所有权从 Worker 转移到 Inventory +claim_pair(): 等待完整 T/Q pair,通过 PairLease 原子转移所有权 +内部单锁 + Condition,不做分片,不做无锁。 +""" +import asyncio +import time +from collections import deque + +from .envelope import ResourceEnvelope + + +class PairLease: + """claim_pair 成功后的 pair 所有权对象。 + + 用法: + async with inventory.claim_pair() as pair: + t_val, q_val = pair.t.value, pair.q.value + ... + 离开上下文时无论成功/失败/取消,都核销 T/Q 并释放 slot。 + """ + + __slots__ = ('t', 'q', '_inventory') + + def __init__(self, t_env, q_env, inventory): + self.t = t_env + self.q = q_env + self._inventory = inventory + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + # 第一版策略: 无论结果如何,均核销 T/Q + self.t.release_slot_once(reason='pair_settle') + self.q.release_slot_once(reason='pair_settle') + # 通知可能在等待的 claim_pair + async with self._inventory._lock: + self._inventory._cond.notify_all() + return False + + +class Inventory: + """T/Q 库存门面。 + + 内部维护两个 deque 和一个 Lock+Condition。 + put_t/put_q 在锁内 O(1) 操作;claim_pair 在锁外等待 Condition。 + """ + + def __init__(self, metrics=None): + self._t_buf = deque() + self._q_buf = deque() + self._lock = asyncio.Lock() + self._cond = asyncio.Condition(self._lock) + self._metrics = metrics + + # ── 入库 ── + + async def put_t(self, env: ResourceEnvelope): + """将 T envelope 入库。调用前 env 属于 Worker,成功后属于 Inventory。""" + async with self._lock: + now = time.time() + if env.is_expired(now): + env.discard(reason='expired_on_put') + if self._metrics: + self._metrics.t_expired += 1 + return + self._cleanup_expired(now) + self._t_buf.append(env) + if self._metrics: + self._metrics.t_admitted += 1 + self._cond.notify_all() + + async def put_q(self, env: ResourceEnvelope): + """将 Q envelope 入库。调用前 env 属于 Worker,成功后属于 Inventory。""" + async with self._lock: + now = time.time() + if env.is_expired(now): + env.discard(reason='expired_on_put') + if self._metrics: + self._metrics.q_expired += 1 + return + self._cleanup_expired(now) + self._q_buf.append(env) + if self._metrics: + self._metrics.q_admitted += 1 + self._cond.notify_all() + + # ── Claim ── + + def claim_pair(self): + """返回 PairLease 的 async context manager。 + + 等待期间可取消,取消不弹出资源。成功后 T/Q 原子弹出并转移给 PairLease。 + """ + return _PairClaim(self, self._metrics) + + # ── 内部 ── + + def _try_pop_pair(self): + """在锁内尝试弹出一对未过期的 T/Q。返回 (t_env, q_env) 或 None。""" + now = time.time() + # 清理过期 + self._cleanup_expired(now) + if self._t_buf and self._q_buf: + t_env = self._t_buf.popleft() + q_env = self._q_buf.popleft() + return t_env, q_env + return None + + def _cleanup_expired(self, now): + """lazy cleanup: 清理已过期的 T/Q 并释放 slot。锁内调用。""" + self._cleanup_buf(self._t_buf, now, 't_expired') + self._cleanup_buf(self._q_buf, now, 'q_expired') + + def _cleanup_buf(self, buf, now, metric_name): + """扫描完整库存,保留未过期实体的原始 FIFO 顺序。""" + kept = deque() + while buf: + env = buf.popleft() + if env.is_expired(now): + env.discard(reason='expired_in_inventory') + if self._metrics: + setattr(self._metrics, metric_name, getattr(self._metrics, metric_name) + 1) + else: + kept.append(env) + buf.extend(kept) + + @property + def t_depth(self): + return len(self._t_buf) + + @property + def q_depth(self): + return len(self._q_buf) + + +class _PairClaim: + """claim_pair() 返回的 async context manager。""" + + __slots__ = ('_inventory', '_pair', '_metrics') + + def __init__(self, inventory, metrics): + self._inventory = inventory + self._pair = None + self._metrics = metrics + + async def __aenter__(self): + inv = self._inventory + async with inv._lock: + # 循环等待直到有 pair 或被取消 + while True: + pair = inv._try_pop_pair() + if pair is not None: + t_env, q_env = pair + self._pair = PairLease(t_env, q_env, inv) + if self._metrics: + self._metrics.pair_claimed += 1 + return self._pair + # 没有 pair,等待 notify + await inv._cond.wait() + + async def __aexit__(self, exc_type, exc, tb): + if self._pair is not None: + # PairLease.__aexit__ 处理核销 + return await self._pair.__aexit__(exc_type, exc, tb) + return False diff --git a/grok_register/core/observer.py b/grok_register/core/observer.py new file mode 100644 index 0000000..61ac69b --- /dev/null +++ b/grok_register/core/observer.py @@ -0,0 +1,215 @@ +""" +Observer — 只读观测层 + +只记录指标、生成日志。不修改 Semaphore,不调度 Worker,不释放资源。 +""" +import time +from collections import deque + + +class Metrics: + """全局指标收集器。所有字段只在 asyncio 单线程中写入,无需加锁。""" + + __slots__ = ( + 'start_time', + 't_produced', 't_admitted', 't_claimed', 't_expired', 't_discarded', + 't_solve_count', 't_solve_seconds', 't_solve_failed', + 'solver_goto_seconds', 'solver_inject_seconds', 'solver_initial_seconds', + 'solver_click_seconds', 'solver_wait_seconds', 'solver_reused_count', + 'solver_visible_frame_count', + 's_physical_count', 's_physical_wait_seconds', 's_physical_hold_seconds', + 'p_physical_count', 'p_physical_wait_seconds', 'p_physical_hold_seconds', + 'c_physical_count', 'c_physical_wait_seconds', 'c_physical_hold_seconds', + 'p_email_create_count', 'p_email_create_seconds', + 'p_page_prepare_count', 'p_page_prepare_seconds', + 'p_send_count', 'p_send_seconds', + 'c_page_acquire_count', 'c_page_acquire_seconds', + 'c_verify_count', 'c_verify_seconds', + 'c_register_count', 'c_register_seconds', + 'c_hot_page_hits', 'c_hot_page_misses', + 'q_sent', 'q_returned', 'q_admitted', 'q_claimed', 'q_expired', 'q_discarded', + 'q_send_batches', 'q_send_batch_items', + 'pair_claimed', 'pair_consumed_ok', 'pair_consumed_fail', + 'success_count', + 'registration_starts', + '_clock', 'started_monotonic', 'recent_success_times', + ) + + def __init__(self, clock=time.monotonic): + self._clock = clock + self.started_monotonic = clock() + self.recent_success_times = deque() + self.start_time = time.time() + # T 生命周期 + self.t_produced = 0 + self.t_admitted = 0 + self.t_claimed = 0 + self.t_expired = 0 + self.t_discarded = 0 + self.t_solve_count = 0 + self.t_solve_seconds = 0.0 + self.t_solve_failed = 0 + self.solver_goto_seconds = 0.0 + self.solver_inject_seconds = 0.0 + self.solver_initial_seconds = 0.0 + self.solver_click_seconds = 0.0 + self.solver_wait_seconds = 0.0 + self.solver_reused_count = 0 + self.solver_visible_frame_count = 0 + self.s_physical_count = 0 + self.s_physical_wait_seconds = 0.0 + self.s_physical_hold_seconds = 0.0 + self.p_physical_count = 0 + self.p_physical_wait_seconds = 0.0 + self.p_physical_hold_seconds = 0.0 + self.c_physical_count = 0 + self.c_physical_wait_seconds = 0.0 + self.c_physical_hold_seconds = 0.0 + self.p_email_create_count = 0 + self.p_email_create_seconds = 0.0 + self.p_page_prepare_count = 0 + self.p_page_prepare_seconds = 0.0 + self.p_send_count = 0 + self.p_send_seconds = 0.0 + self.c_page_acquire_count = 0 + self.c_page_acquire_seconds = 0.0 + self.c_verify_count = 0 + self.c_verify_seconds = 0.0 + self.c_register_count = 0 + self.c_register_seconds = 0.0 + self.c_hot_page_hits = 0 + self.c_hot_page_misses = 0 + # Q 生命周期 + self.q_sent = 0 + self.q_returned = 0 + self.q_admitted = 0 + self.q_claimed = 0 + self.q_expired = 0 + self.q_discarded = 0 + self.q_send_batches = 0 + self.q_send_batch_items = 0 + # Pair + self.pair_claimed = 0 + self.pair_consumed_ok = 0 + self.pair_consumed_fail = 0 + # 成功数 + self.success_count = 0 + self.registration_starts = 0 + + def next_registration_task(self): + self.registration_starts += 1 + return self.registration_starts + + def record_success(self): + self.success_count += 1 + self.recent_success_times.append(self._clock()) + + def five_minute_success_rate(self): + now = self._clock() + cutoff = now - 300.0 + while self.recent_success_times and self.recent_success_times[0] < cutoff: + self.recent_success_times.popleft() + if not self.recent_success_times: + return None if self.success_count == 0 else 0.0 + elapsed = max(1.0, min(300.0, now - self.started_monotonic)) + return len(self.recent_success_times) * 60.0 / elapsed + + def runtime_average_success_rate(self): + if self.success_count == 0: + return None + elapsed = max(1.0, self._clock() - self.started_monotonic) + return self.success_count * 60.0 / elapsed + + def snapshot(self, inventory, sems): + """生成一行监控日志。""" + elapsed = time.time() - self.start_time + rate = self.success_count / (elapsed / 60) if elapsed > 60 else 0 + p_batch_avg = ( + self.q_send_batch_items / self.q_send_batches + if self.q_send_batches else 0 + ) + t_solve_avg = ( + self.t_solve_seconds / self.t_solve_count + if self.t_solve_count else 0 + ) + solver_goto_avg = ( + self.solver_goto_seconds / self.t_solve_count + if self.t_solve_count else 0 + ) + solver_inject_avg = ( + self.solver_inject_seconds / self.t_solve_count + if self.t_solve_count else 0 + ) + solver_initial_avg = ( + self.solver_initial_seconds / self.t_solve_count + if self.t_solve_count else 0 + ) + solver_click_avg = ( + self.solver_click_seconds / self.t_solve_count + if self.t_solve_count else 0 + ) + solver_wait_avg = ( + self.solver_wait_seconds / self.t_solve_count + if self.t_solve_count else 0 + ) + solver_reuse_ratio = ( + self.solver_reused_count / self.t_solve_count + if self.t_solve_count else 0 + ) + solver_visible_ratio = ( + self.solver_visible_frame_count / self.t_solve_count + if self.t_solve_count else 0 + ) + s_phys_wait, s_phys_hold = self._avg_pair( + self.s_physical_wait_seconds, self.s_physical_hold_seconds, self.s_physical_count + ) + p_phys_wait, p_phys_hold = self._avg_pair( + self.p_physical_wait_seconds, self.p_physical_hold_seconds, self.p_physical_count + ) + c_phys_wait, c_phys_hold = self._avg_pair( + self.c_physical_wait_seconds, self.c_physical_hold_seconds, self.c_physical_count + ) + p_email_create = self._avg(self.p_email_create_seconds, self.p_email_create_count) + p_page_prepare = self._avg(self.p_page_prepare_seconds, self.p_page_prepare_count) + p_send_stage = self._avg(self.p_send_seconds, self.p_send_count) + c_page_acquire = self._avg(self.c_page_acquire_seconds, self.c_page_acquire_count) + c_verify = self._avg(self.c_verify_seconds, self.c_verify_count) + c_register = self._avg(self.c_register_seconds, self.c_register_count) + p_send_sem = sems.get("p_send") + admission = sems.get("admission") + p_send_part = f' p_send:{p_send_sem._value}' if p_send_sem is not None else '' + admission_part = ( + f' t_prog:{admission.t_in_progress} q_inflight:{admission.q_inflight}' + if admission is not None else '' + ) + return ( + f'[*] T:{inventory.t_depth} Q:{inventory.q_depth} ' + f'phys:{sems["physical"]._value}{p_send_part} t_slot:{sems["t_slot"]._value} ' + f'q_slot:{sems["q_slot"]._value} q_pend:{sems["q_pending"]._value} ' + f'p_batch:{p_batch_avg:.1f}{admission_part} ' + f's_phys:{s_phys_wait:.2f}/{s_phys_hold:.2f} ' + f'p_phys:{p_phys_wait:.2f}/{p_phys_hold:.2f} ' + f'c_phys:{c_phys_wait:.2f}/{c_phys_hold:.2f} ' + f'p_stage:{p_email_create:.2f}/{p_page_prepare:.2f}/{p_send_stage:.2f} ' + f'c_stage:{c_page_acquire:.2f}/{c_verify:.2f}/{c_register:.2f} ' + f'c_hot:{self.c_hot_page_hits}/{self.c_hot_page_misses} ' + f't_solve_avg:{t_solve_avg:.1f} t_solve_fail:{self.t_solve_failed} ' + f'solver_goto:{solver_goto_avg:.2f} solver_inject:{solver_inject_avg:.2f} ' + f'solver_initial:{solver_initial_avg:.2f} solver_click:{solver_click_avg:.2f} ' + f'solver_wait:{solver_wait_avg:.2f} solver_reuse:{solver_reuse_ratio:.2f} ' + f'solver_visible:{solver_visible_ratio:.2f} ' + f't_prod:{self.t_produced} t_adm:{self.t_admitted} t_exp:{self.t_expired} ' + f'q_sent:{self.q_sent} q_ret:{self.q_returned} q_adm:{self.q_admitted} q_exp:{self.q_expired} ' + f'pair:{self.pair_claimed} ok:{self.pair_consumed_ok} fail:{self.pair_consumed_fail} ' + f'rate:{rate:.1f}/min #{self.success_count}' + ) + + @staticmethod + def _avg(total, count): + return total / count if count else 0 + + @staticmethod + def _avg_pair(wait_total, hold_total, count): + if not count: + return 0, 0 + return wait_total / count, hold_total / count diff --git a/grok_register/email_server.py b/grok_register/email_server.py new file mode 100644 index 0000000..6dc2099 --- /dev/null +++ b/grok_register/email_server.py @@ -0,0 +1,171 @@ +""" +邮件接收 API 服务器 +==================== +接收 Cloudflare Email Routing 转发的邮件,供注册服务查询。 + +端点: + POST /webhook — Cloudflare 转发邮件到这里 + GET /check/ — 查询某邮箱的验证码 + GET /domains — 返回可用域名 + GET /health — 健康检查 + +用法: + bash start.sh --email-service + EMAIL_DOMAIN=your.domain bash start.sh --email-service --port 8080 +""" +import os, re, json, time, sys +try: + from dotenv import load_dotenv + load_dotenv() +except ImportError: + pass +from http.server import HTTPServer, ThreadingHTTPServer, BaseHTTPRequestHandler +from urllib.parse import urlparse, parse_qs +from threading import Lock + +# 配置 +DEFAULT_DOMAIN = os.environ.get("EMAIL_DOMAIN", "") +DEFAULT_PORT = 8080 + +# 存储 +emails = {} # {email_address: [{"code": "ABC123", "time": timestamp, "raw": "..."}]} +emails_lock = Lock() + +# 清理过期邮件(5 分钟) +def cleanup_old(): + now = time.time() + with emails_lock: + for addr in list(emails.keys()): + emails[addr] = [e for e in emails[addr] if now - e['time'] < 300] + if not emails[addr]: + del emails[addr] + + +def extract_code(text): + """从邮件内容提取验证码(ABC-DEF 或 ABCDEF 格式)""" + # 格式1: ABC-DEF + m = re.search(r'>([A-Z0-9]{3}-[A-Z0-9]{3})<', text) + if m: + return m.group(1).replace('-', '') + # 格式2: 直接 6 位 + m = re.search(r'>([A-Z0-9]{6})<', text) + if m: + return m.group(1) + # 格式3: 正文中的 6 位 + m = re.search(r'\b([A-Z0-9]{3}-?[A-Z0-9]{3})\b', text) + if m: + return m.group(1).replace('-', '') + return None + + +class EmailHandler(BaseHTTPRequestHandler): + def do_POST(self): + if self.path == '/webhook': + content_length = int(self.headers.get('Content-Length', 0)) + body = self.rfile.read(content_length).decode('utf-8') + + try: + data = json.loads(body) + except: + data = {} + + # Cloudflare Email Routing 格式 + to_addr = data.get('to', data.get('recipient', '')) + from_addr = data.get('from', data.get('sender', '')) + subject = data.get('subject', '') + text = data.get('text', '') + html = data.get('html', '') + + # 提取验证码 + content = f"{subject}\n{text}\n{html}" + code = extract_code(content) + + if to_addr and code: + with emails_lock: + if to_addr not in emails: + emails[to_addr] = [] + emails[to_addr].append({ + 'code': code, + 'time': time.time(), + 'from': from_addr, + 'subject': subject + }) + print(f'[+] {to_addr} code={code}', flush=True) + + self.send_response(200) + self.send_header('Content-Type', 'application/json') + self.end_headers() + self.wfile.write(json.dumps({"ok": True, "code": code}).encode()) + else: + self.send_response(404) + self.end_headers() + + def do_GET(self): + parsed = urlparse(self.path) + path = parsed.path + + if path == '/health': + self._json({"status": "ok", "emails": len(emails)}) + + elif path == '/domains': + self._json({"domains": [DEFAULT_DOMAIN]}) + + elif path.startswith('/check/'): + addr = path[7:] # 去掉 /check/ + cleanup_old() + with emails_lock: + items = emails.get(addr, []) + if items: + # 返回最新的验证码 + latest = items[-1] + self._json({"code": latest['code'], "from": latest['from']}) + else: + self._json({"code": None}) + + elif path == '/list': + # 列出所有有邮件的地址(调试用) + cleanup_old() + with emails_lock: + result = {addr: len(msgs) for addr, msgs in emails.items()} + self._json(result) + + else: + self.send_response(404) + self.end_headers() + + def _json(self, data): + body = json.dumps(data).encode() + self.send_response(200) + self.send_header('Content-Type', 'application/json') + self.send_header('Content-Length', str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, format, *args): + print(f"[HTTP] {args[0] if args else format}", flush=True) + + +def main(): + global DEFAULT_DOMAIN + port = DEFAULT_PORT + domain = DEFAULT_DOMAIN + + for i, arg in enumerate(sys.argv[1:]): + if arg == '--port' and i + 2 <= len(sys.argv): + port = int(sys.argv[i + 2]) + if arg == '--domain' and i + 2 <= len(sys.argv): + domain = sys.argv[i + 2] + + DEFAULT_DOMAIN = domain + + print(f"[*] Email server starting on :{port}", flush=True) + print(f"[*] Domain: {domain}", flush=True) + print(f"[*] Webhook: http://0.0.0.0:{port}/webhook", flush=True) + print(f"[*] Check: http://localhost:{port}/check/", flush=True) + + server = ThreadingHTTPServer(('0.0.0.0', port), EmailHandler) + server.serve_forever() + + +if __name__ == "__main__": + main() diff --git a/grok_register/register.py b/grok_register/register.py new file mode 100644 index 0000000..26b5633 --- /dev/null +++ b/grok_register/register.py @@ -0,0 +1,2308 @@ +""" +Grok Free Register — CSP 异步并发架构 +============================================= +单进程 asyncio + 单共享 CloakBrowser + Semaphore 背压: + - S_Worker: 生成 Turnstile token (T) + - P_Worker: 创建邮箱 + 发送验证码 + 轮询验证码 (Q) + - C_Worker: claim pair 并执行注册 + - Semaphore 背压控制容量,无需中心调度器 + +两种邮箱模式(EMAIL_MODE): + - tempmail (默认,零配置): 免费临时邮箱,多 provider 自动 fallback + - custom: 自建域名邮箱,Cloudflare Email Routing → Worker → 本地 webhook + (见 grok_register/email_server.py / cloudflare/email-worker.js) + +配置全部走环境变量 / .env(见 .env.example);CLI: --max-mem 6G --target 100 +用法: + bash start.sh # 一键引导 +""" +import os, json, random, string, time, re, secrets, base64, struct, asyncio, glob, sys, multiprocessing +from contextlib import suppress +try: + from dotenv import load_dotenv + load_dotenv() +except ImportError: + pass +import requests as req +from urllib.parse import quote +from playwright.async_api import async_playwright +from concurrent.futures import ThreadPoolExecutor + +# CSP 架构组件 +from grok_register.core.admission import AdmissionGate +from grok_register.core.envelope import ResourceEnvelope +from grok_register.core.inventory import Inventory +from grok_register.core.observer import Metrics +from grok_register.clearance import ( + apply_clearance_to_context, + cached_user_agent, + format_prewarm_log, + playwright_proxy_settings, + prewarm_clearance, + register_proxy_url, +) + +os.makedirs("keys", exist_ok=True) +SITE_URL = "https://accounts.x.ai" + +# ── 配置(环境变量 / .env,见 .env.example)── +def _env_int(key, default): + try: + return int(str(os.environ.get(key, "")).strip() or default) + except ValueError: + return default + +def _env_int_or_none(key): + raw = str(os.environ.get(key, "")).strip() + if not raw: + return None + try: + return int(raw) + except ValueError: + return None + +EMAIL_MODE = (os.environ.get("EMAIL_MODE") or "tempmail").strip().lower() # tempmail | custom +if EMAIL_MODE == "mailtm": # 兼容旧名 + EMAIL_MODE = "tempmail" +LOCAL_EMAIL_API = (os.environ.get("EMAIL_API") or "http://127.0.0.1:8080").strip() +EMAIL_DOMAIN = (os.environ.get("EMAIL_DOMAIN") or "").strip() +MIN_FREE_MEM_MB = _env_int("MIN_FREE_MEM_MB", 500) # 自动容量派生时保留的内存(MB) +T_TARGET = _env_int("T_TARGET", 4) # token 池缓冲目标 +Q_TARGET = _env_int("Q_TARGET", 4) # 就绪验证码缓冲目标 +TARGET = _env_int("TARGET", 0) # 攒够 N 个号自动停(0=不限;--target N 可覆盖) + +# CSP 容量参数 +PHYSICAL_CAP = _env_int("PHYSICAL_CAP", 0) # 本地物理资源许可,0=自动派生 +PHYSICAL_PER_CPU = _env_int("PHYSICAL_PER_CPU", 2) # 自动派生 CPU 侧保守上限;压测可临时覆盖 +PHYSICAL_MEM_MB = _env_int("PHYSICAL_MEM_MB", 512) # 每个物理许可的保守内存预算(MB) +CAPACITY_PROFILE = (os.environ.get("CAPACITY_PROFILE") or "").strip() +T_SLOT_CAP = _env_int("T_SLOT_CAP", 8) # token 库存缓冲 +Q_SLOT_CAP = _env_int("Q_SLOT_CAP", 8) # 验证码库存缓冲 +Q_PENDING_CAP = _env_int("Q_PENDING_CAP", 12) # 外部在途 Q 请求上限 +T_MAX_AGE = _env_int("T_MAX_AGE", 300) # token 最大年龄(秒) +Q_MAX_AGE = _env_int("Q_MAX_AGE", 120) # 验证码最大年龄(秒) +P_REQUEST_TIMEOUT = _env_int("P_REQUEST_TIMEOUT", 95) # P 等待 Q 返回超时(秒) +C_CONSUME_TIMEOUT = _env_int("C_CONSUME_TIMEOUT", 60) # C 消费完整 pair 超时(秒) +S_WORKERS = _env_int("S_WORKERS", 0) # 0=自动 +P_WORKERS = _env_int("P_WORKERS", 0) +C_WORKERS = _env_int("C_WORKERS", 0) +C_HOT_PAGE_POOL = (os.environ.get("C_HOT_PAGE_POOL", "0").strip().lower() in ("1", "true", "yes")) +C_HOT_PAGE_POOL_SIZE = _env_int("C_HOT_PAGE_POOL_SIZE", 0) +C_SET_COOKIE_VIA_REQUEST = ( + os.environ.get("C_SET_COOKIE_VIA_REQUEST", "1" if C_HOT_PAGE_POOL else "0") + .strip() + .lower() + in ("1", "true", "yes") +) + +# CSP v2 局部门控/批量发送参数。水位默认在启动期结合 Physical_Sem 派生。 +_T_HIGH_WATER_OVERRIDE = _env_int_or_none("T_HIGH_WATER") +_T_LOW_WATER_OVERRIDE = _env_int_or_none("T_LOW_WATER") +_Q_HIGH_WATER_OVERRIDE = _env_int_or_none("Q_HIGH_WATER") +_Q_LOW_WATER_OVERRIDE = _env_int_or_none("Q_LOW_WATER") +P_BATCH_MAX = max(1, _env_int("P_BATCH_MAX", 4)) +P_SEND_CAP = _env_int("P_SEND_CAP", 0) # >0=显式限制并发 P 发送页面;0=不额外建模 +PAGE_GOTO_WAIT_UNTIL = os.environ.get("PAGE_GOTO_WAIT_UNTIL", "domcontentloaded").strip() or "domcontentloaded" +PAGE_POST_WAIT_MS = _env_int("PAGE_POST_WAIT_MS", 500) +PAGE_BLOCK_STATIC_ASSETS = ( + os.environ.get("PAGE_BLOCK_STATIC_ASSETS", "0").strip().lower() + in ("1", "true", "yes") +) +REGISTRATION_DIAGNOSTICS = ( + os.environ.get("REGISTRATION_DIAGNOSTICS", "0").strip().lower() + in ("1", "true", "yes") +) +REGISTRATION_RATE_LIMIT_COOLDOWN = max( + 60, _env_int("REGISTRATION_RATE_LIMIT_COOLDOWN", 60) +) +REGISTRATION_RATE_LIMIT_RECOVERY_SECONDS = max( + 1, _env_int("REGISTRATION_RATE_LIMIT_RECOVERY_SECONDS", 60) +) +REGISTRATION_RATE_LIMIT_RECOVERY_INTERVAL = max( + 1, _env_int("REGISTRATION_RATE_LIMIT_RECOVERY_INTERVAL", 3) +) +REGISTER_LOG_MODE = "user" + +SITE_KEY = None +ACTION_ID = None +STATE_TREE = None + +start_time = time.time() +success_count = 0 +file_lock = asyncio.Lock() +STOP = asyncio.Event() + +# 角色标识 + 轮询/HTTP 专用线程池(与 CPU 密集的浏览器操作解耦) +SOLVE, PRODUCE, CONSUME, IDLE = 'SOLVE', 'PRODUCE', 'CONSUME', 'IDLE' +POLL_EXECUTOR = ThreadPoolExecutor(max_workers=32) + +def resolve_register_log_mode(argv=None, env=None): + argv = list([] if argv is None else argv) + env = dict(os.environ if env is None else env) + mode = (env.get("REGISTER_LOG_MODE") or "user").strip().lower() + if "--debug" in argv: + mode = "debug" + if mode not in {"user", "debug"}: + raise ValueError("REGISTER_LOG_MODE must be user or debug") + return mode + + +def _terminal_output(msg): + print(msg, flush=True) + + +def log(msg): + try: + _terminal_output(msg) + except Exception: + return + + +def debug_log(msg): + if REGISTER_LOG_MODE == "debug": + log(msg) +def rand_str(n=15): return ''.join(random.choice(string.ascii_lowercase + string.digits) for _ in range(n)) + + +def sanitize_terminal_error(error): + return type(error).__name__ + + +def format_user_registration_event( + kind, + *, + task_id=None, + count=None, + rate_per_minute=None, + wait_seconds=None, + remaining=None, +): + if kind == "service_started": + progress = f"剩余 {remaining}" if remaining is not None else "持续运行" + return f"[✓] 注册服务已启动 | {progress}" + if kind == "started": + suffix = f" | 剩余 {remaining}" if remaining is not None else "" + return f"[→] 开始注册 #{task_id}{suffix}" + if kind == "success": + rate = "—" if rate_per_minute is None else f"{rate_per_minute:.1f}/分" + return f"[✓] 注册成功 #{task_id} | 运行平均 {rate} | 累计 {count}" + if kind == "failed": + return f"[✗] 注册失败 #{task_id} | 已跳过,继续下一任务" + if kind == "rate_limited": + return f"[⏸] 触发限流 | {wait_seconds}秒后恢复探测" + if kind == "recovered": + return f"[▶] 限流解除 | 实际等待 {wait_seconds}秒" + if kind == "stopped": + return f"[■] 注册服务已停止 | 累计 {count or 0}" + raise ValueError(f"unknown user registration event: {kind}") + + +class RegistrationRateLimited(RuntimeError): + """注册提交被目标站点的限流页替代。""" + + +class RegistrationRateLimitCircuit: + """在检测到注册限流后暂停新的 C 阶段提交。""" + + def __init__( + self, + cooldown_seconds, + recovery_seconds=60, + recovery_interval=3, + clock=time.monotonic, + ): + self.cooldown_seconds = cooldown_seconds + self.recovery_seconds = recovery_seconds + self.recovery_interval = recovery_interval + self._clock = clock + self._blocked_until = 0.0 + self._tripped_at = None + self._probe_active = False + self._probe_token = None + self._recovering_until = 0.0 + self._next_recovery_submit = 0.0 + + def remaining_seconds(self): + return max(0, int(self._blocked_until - self._clock() + 0.999)) + + def is_open(self): + return self.remaining_seconds() > 0 + + def trip(self): + starts_new_window = not self.is_open() + if self._tripped_at is None: + self._tripped_at = self._clock() + self._recovering_until = 0.0 + self._next_recovery_submit = 0.0 + self._probe_active = False + self._probe_token = None + self._blocked_until = max( + self._blocked_until, + self._clock() + self.cooldown_seconds, + ) + return starts_new_window + + async def wait(self): + while True: + if self.is_open(): + await asyncio.sleep(min(self.remaining_seconds(), 5)) + continue + if self._tripped_at is None: + if not self._recovering_until: + return False + if self._probe_active: + await asyncio.sleep(0.2) + continue + if self._clock() >= self._recovering_until: + self._recovering_until = 0.0 + self._next_recovery_submit = 0.0 + return False + recovery_wait = self._next_recovery_submit - self._clock() + if recovery_wait > 0: + await asyncio.sleep(min(recovery_wait, 0.5)) + continue + if not self._probe_active: + self._probe_active = True + self._probe_token = object() + return self._probe_token + await asyncio.sleep(0.2) + + def can_submit(self, probe_token=False): + """仅允许正常态任务或当前恢复探针发起注册提交。""" + if self._tripped_at is None and not self._recovering_until: + return True + return ( + probe_token is not False + and probe_token is self._probe_token + and not self.is_open() + ) + + def consume_recovery_seconds(self, probe_token=None): + if self._tripped_at is None: + self.complete_recovery_submission(probe_token) + return None + if probe_token is not None and probe_token is not self._probe_token: + return None + elapsed = self._clock() - self._tripped_at + self._tripped_at = None + self._blocked_until = 0.0 + self._probe_active = False + self._probe_token = None + self._recovering_until = self._clock() + self.recovery_seconds + self._next_recovery_submit = self._clock() + self.recovery_interval + return elapsed + + def complete_recovery_submission(self, probe_token): + """完成恢复期内的一次成功提交,并按固定节奏放行下一项。""" + if probe_token is not self._probe_token or not self._recovering_until: + return False + self._probe_active = False + self._probe_token = None + self._next_recovery_submit = self._clock() + self.recovery_interval + return True + + def release_probe(self, probe_token): + """提交前资源已失效时让出探针,不额外增加冷却窗口。""" + if probe_token is self._probe_token: + self._probe_active = False + self._probe_token = None + + def defer_probe(self, probe_token): + """真正的探针失败后重新进入完整冷却。""" + if probe_token is not self._probe_token: + return + self._probe_active = False + self._probe_token = None + if self._tripped_at is not None: + self._blocked_until = max( + self._blocked_until, + self._clock() + self.cooldown_seconds, + ) + elif self._recovering_until: + self._next_recovery_submit = self._clock() + self.recovery_interval + + +REGISTRATION_RATE_LIMIT_CIRCUIT = RegistrationRateLimitCircuit( + REGISTRATION_RATE_LIMIT_COOLDOWN, + recovery_seconds=REGISTRATION_RATE_LIMIT_RECOVERY_SECONDS, + recovery_interval=REGISTRATION_RATE_LIMIT_RECOVERY_INTERVAL, +) + +def _signup_response_markers(text): + """将失败响应归类为固定标签,诊断时不输出任何服务端正文。""" + lowered = text.lower() + groups = { + "challenge": ("captcha", "cf-chl", "challenge-platform"), + "rate_limited": ("rate limit", "too many requests", "try again later"), + "signin_page": ("sign in", "log in"), + "signup_page": ("sign up", "create your account"), + "next_page": ("__next", "/_next/"), + "action_error": ("server action", "next-action", "digest"), + } + return ",".join(name for name, needles in groups.items() if any(x in lowered for x in needles)) or "unclassified" + +def pb_varint(n): + parts = [] + while n > 0x7f: parts.append((n & 0x7f) | 0x80); n >>= 7 + parts.append(n); return bytes(parts) +def pb_str(fid, val): + vb = val.encode() + return struct.pack('B', (fid << 3) | 2) + pb_varint(len(vb)) + vb +def decode_jwt_payload(token): + parts = token.split('.') + if len(parts) < 2: return None + payload = parts[1] + '=' * (4 - len(parts[1]) % 4) + try: return json.loads(base64.urlsafe_b64decode(payload)) + except Exception: return None +def find_chrome(): + paths = glob.glob(os.path.expanduser("~/.cloakbrowser/chromium-*/chrome")) + if not paths: raise RuntimeError("CloakBrowser not found") + return sorted(paths)[-1] + + +def _read_cmdline(pid: int) -> str: + try: + with open(f"/proc/{pid}/cmdline", "rb") as fh: + return fh.read().replace(b"\x00", b" ").decode("utf-8", "replace") + except (FileNotFoundError, ProcessLookupError, PermissionError, OSError): + return "" + + +def _is_cloakbrowser_chrome(cmdline: str) -> bool: + """注册机用的 CloakBrowser 路径;排除 docker FS 的 /usr/lib/chromium。""" + if not cmdline: + return False + if "cloakbrowser" in cmdline: + return True + # playwright 偶发只带 chrome 绝对路径 + if "/.cloakbrowser/" in cmdline and "chrome" in cmdline: + return True + return False + + +def reap_orphan_register_browsers(*, exclude_pids: set[int] | None = None) -> int: + """杀掉无主的 CloakBrowser/chrome(上次 Ctrl-C 失败或 kill -9 残留)。 + + 只动 host 上 cloakbrowser 路径的进程,不动 docker 里 FlareSolverr 的 chromium。 + """ + import signal as _signal + import time as _time + + exclude = set(exclude_pids or ()) + exclude.add(os.getpid()) + # 当前进程树:子进程先别杀,留给 browser.close + try: + for line in os.popen(f"ps -o pid= --ppid {os.getpid()}").read().split(): + try: + exclude.add(int(line)) + except ValueError: + pass + except Exception: + pass + + victims: list[int] = [] + try: + for name in os.listdir("/proc"): + if not name.isdigit(): + continue + pid = int(name) + if pid in exclude or pid <= 1: + continue + cmd = _read_cmdline(pid) + if not _is_cloakbrowser_chrome(cmd): + continue + # 跳过 docker cgroup(双保险) + try: + with open(f"/proc/{pid}/cgroup", "r", encoding="utf-8") as fh: + if "docker" in fh.read(): + continue + except Exception: + pass + victims.append(pid) + except Exception: + return 0 + + if not victims: + return 0 + + for pid in victims: + try: + os.kill(pid, _signal.SIGTERM) + except ProcessLookupError: + pass + except PermissionError: + pass + _time.sleep(0.4) + for pid in victims: + try: + os.kill(pid, 0) + except ProcessLookupError: + continue + try: + os.kill(pid, _signal.SIGKILL) + except ProcessLookupError: + pass + except PermissionError: + pass + return len(victims) + + +async def _close_browser_hard(browser, timeout: float = 5.0) -> None: + """先正常 close;超时则按进程树 SIGKILL,避免重启叠浏览器。""" + if browser is None: + return + proc = getattr(browser, "process", None) + pid = None + if proc is not None: + pid = getattr(proc, "pid", None) + try: + await asyncio.wait_for(browser.close(), timeout=timeout) + return + except Exception: + pass + if pid: + try: + # 杀进程组(playwright 常把 chrome 放同组) + os.killpg(os.getpgid(pid), 9) + except Exception: + try: + os.kill(pid, 9) + except Exception: + pass + # 扫一遍残留 cloakbrowser + await asyncio.to_thread(reap_orphan_register_browsers) + + +# ────────────────────────────────────────────── +# 资源检测 +# ────────────────────────────────────────────── +def get_system_resources(max_mem_arg=None): + import subprocess + cpu_count = multiprocessing.cpu_count() + try: + out = subprocess.check_output(["free", "-m"]).decode() + for line in out.split("\n"): + if "Mem" in line: + parts = line.split() + total, available = int(parts[1]), int(parts[6]) + break + else: + total, available = 4096, 2048 + except Exception: + total, available = 4096, 2048 + + if max_mem_arg: + if max_mem_arg.endswith('%'): + max_mem = int(total * float(max_mem_arg[:-1]) / 100) + elif max_mem_arg.upper().endswith('G'): + max_mem = int(float(max_mem_arg[:-1]) * 1024) + elif max_mem_arg.upper().endswith('M'): + max_mem = int(max_mem_arg[:-1]) + else: + max_mem = int(max_mem_arg) + else: + max_mem = available + + return {'cpu': cpu_count, 'total_mem': total, 'available_mem': available, 'max_mem': max_mem} + + +def load_capacity_profile(path=CAPACITY_PROFILE): + """读取离线校准生成的设备 profile。不存在或无效时返回空配置。""" + if not path: + return {} + try: + with open(path, "r") as f: + data = json.load(f) + except Exception: + return {} + if not isinstance(data, dict): + return {} + profile = {} + try: + physical_cap = int(data.get("physical_cap", 0)) + except (TypeError, ValueError): + physical_cap = 0 + if physical_cap > 0: + profile["physical_cap"] = physical_cap + return profile + + +def derive_capacity( + cpu_count, + max_mem_mb, + *, + physical_cap=None, + profile_physical_cap=None, + physical_per_cpu=None, + physical_mem_mb=None, + min_free_mem_mb=None, +): + """启动期静态容量派生:显式配置 > 设备 profile > CPU/内存保守自动值。""" + configured_physical = PHYSICAL_CAP if physical_cap is None else physical_cap + profiled_physical = profile_physical_cap or 0 + per_cpu = PHYSICAL_PER_CPU if physical_per_cpu is None else physical_per_cpu + mem_per_physical = PHYSICAL_MEM_MB if physical_mem_mb is None else physical_mem_mb + reserve_mem = MIN_FREE_MEM_MB if min_free_mem_mb is None else min_free_mem_mb + + cpu_cap = max(1, cpu_count * max(1, per_cpu)) + usable_mem = max(0, max_mem_mb - reserve_mem) + mem_cap = max(1, usable_mem // max(1, mem_per_physical)) + auto_cap = max(1, min(cpu_cap, mem_cap)) + + if configured_physical > 0: + physical = configured_physical + elif profiled_physical > 0: + physical = max(1, min(profiled_physical, mem_cap)) + else: + physical = auto_cap + + s_workers = S_WORKERS if S_WORKERS > 0 else physical + 2 + p_workers = P_WORKERS if P_WORKERS > 0 else Q_PENDING_CAP + 2 + c_workers = C_WORKERS if C_WORKERS > 0 else physical + 2 + return physical, s_workers, p_workers, c_workers + + +def derive_admission_watermarks( + physical_cap, + *, + t_slot_cap=None, + q_pending_cap=None, + t_target=None, + q_target=None, + t_high_override=None, + t_low_override=None, + q_high_override=None, + q_low_override=None, +): + """派生 CSP v2 局部门控水位。 + + T 的默认高水位跟随 Physical_Sem,避免物理并发提高后仍只允许少量 T + in-progress。显式环境变量覆盖仍然优先。 + """ + t_slot = T_SLOT_CAP if t_slot_cap is None else t_slot_cap + q_pending = Q_PENDING_CAP if q_pending_cap is None else q_pending_cap + t_goal = T_TARGET if t_target is None else t_target + q_goal = Q_TARGET if q_target is None else q_target + + t_high_cfg = _T_HIGH_WATER_OVERRIDE if t_high_override is None else t_high_override + t_low_cfg = _T_LOW_WATER_OVERRIDE if t_low_override is None else t_low_override + q_high_cfg = _Q_HIGH_WATER_OVERRIDE if q_high_override is None else q_high_override + q_low_cfg = _Q_LOW_WATER_OVERRIDE if q_low_override is None else q_low_override + + if t_high_cfg is None: + t_high = min(max(1, t_slot), max(1, t_goal, physical_cap)) + else: + t_high = max(1, min(max(1, t_slot), t_high_cfg)) + + if t_low_cfg is None: + t_low = max(0, min(t_high, t_high // 2)) + else: + t_low = max(0, min(t_high, t_low_cfg)) + + if q_high_cfg is None: + q_high = max(1, q_pending) + else: + q_high = max(1, min(max(1, q_pending), q_high_cfg)) + + if q_low_cfg is None: + q_low = max(0, min(q_high, q_goal, q_high // 2)) + else: + q_low = max(0, min(q_high, q_low_cfg)) + + return { + "t_low": t_low, + "t_high": t_high, + "q_low": q_low, + "q_high": q_high, + } + + +def derive_c_hot_page_pool_size(physical_cap, c_workers, configured_size=None): + """启动期静态派生 C 热页池容量;显式配置优先。""" + configured = C_HOT_PAGE_POOL_SIZE if configured_size is None else configured_size + if configured and configured > 0: + return configured + return max(1, min(max(1, physical_cap), max(1, c_workers))) + + +# ────────────────────────────────────────────── +# 配置获取 +# ────────────────────────────────────────────── +def _launch_options(): + """Playwright chromium.launch kwargs. HTTP_PROXY/REGISTER_PROXY apply here (not to FlareSolverr).""" + options = {"executable_path": find_chrome(), "headless": True} + proxy = playwright_proxy_settings() + if proxy: + options["proxy"] = proxy + return options + + +async def _prepare_browser_page(browser): + """new_page + inject prewarmed CF cookies into its context when available.""" + page = await browser.new_page() + await apply_clearance_to_context(page.context) + return page + + +async def _prepare_browser_context(browser): + """new_context with optional FS user-agent + CF cookies.""" + kwargs = {} + ua = cached_user_agent() + if ua: + kwargs["user_agent"] = ua + context = await browser.new_context(**kwargs) + await apply_clearance_to_context(context) + return context + + +async def fetch_config(): + global SITE_KEY, ACTION_ID, STATE_TREE + debug_log('[*] Fetching config...') + async with async_playwright() as pw: + browser = await pw.chromium.launch(**_launch_options()) + try: + page = await _prepare_browser_page(browser) + await page.goto(f'{SITE_URL}/sign-up?redirect=grok-com', timeout=30000) + await page.wait_for_timeout(5000) + html = await page.content() + m = re.search(r'0x4AAAAAAA[a-zA-Z0-9_-]+', html) + if m: SITE_KEY = m.group(0); debug_log(f'[+] SITE_KEY: {SITE_KEY}') + for chunk in re.findall(r'self\.__next_f\.push\(\[1,"(.*?)"\]\)', html, re.DOTALL): + if 'sign-up' not in chunk: continue + decoded = chunk.replace('\\"', '"') + f_match = re.search(r'"f":\[\[\[', decoded) + if not f_match: continue + f_start = f_match.start() + 5 + end_idx = decoded.find('"$undefined"', f_start) + if end_idx < 0: continue + STATE_TREE = quote(decoded[f_start:end_idx].replace('\\\\"', '"').replace('\\', ''), safe='') + debug_log(f'[+] STATE_TREE: {STATE_TREE[:50]}...') + break + js_urls = re.findall(r'src="(/_next/static/[^"]+\.js)"', html) + for js_url in js_urls[:50]: + try: + js = await page.evaluate(f"(async()=>{{return await fetch('{js_url}').then(r=>r.text()).catch(()=>\"\" )}})()") + if not js: continue + if not any(kw in js for kw in ['createUser','registerUser','emailValidation']): continue + hexes = re.findall(r'[a-fA-F0-9]{40,50}', js) + if hexes: ACTION_ID = hexes[0]; break + except asyncio.CancelledError: + raise + except Exception: + continue + if ACTION_ID: debug_log(f'[+] ACTION_ID: {ACTION_ID}') + finally: + await _close_browser_hard(browser, timeout=5.0) + if not all([SITE_KEY, ACTION_ID, STATE_TREE]): + raise RuntimeError("Config fetch failed") + + +# ────────────────────────────────────────────── +# 异步操作 +# ────────────────────────────────────────────── +async def grpc_create_code(page, email): + inner = pb_str(1, email) + frame = b'\x00' + struct.pack('>I', len(inner)) + inner + fb64 = base64.b64encode(frame).decode() + s = await page.evaluate(f"(async()=>{{var fb=Uint8Array.from(atob('{fb64}'),c=>c.charCodeAt(0));var r=await fetch('{SITE_URL}/auth_mgmt.AuthManagement/CreateEmailValidationCode',{{method:'POST',headers:{{'content-type':'application/grpc-web+proto','x-grpc-web':'1','x-user-agent':'connect-es/2.1.1'}},body:fb.buffer}});return r.headers.get('grpc-status')||'0';}})()") + return s == '0' + + +async def _prepare_signup_page(page, *, redirect=True, timeout=30000): + if PAGE_BLOCK_STATIC_ASSETS: + await page.route("**/*", _route_static_asset_filter) + url = f'{SITE_URL}/sign-up?redirect=grok-com' if redirect else f'{SITE_URL}/sign-up' + await page.goto(url, timeout=timeout, wait_until=PAGE_GOTO_WAIT_UNTIL) + if PAGE_POST_WAIT_MS > 0: + await page.wait_for_timeout(PAGE_POST_WAIT_MS) + + +async def _route_static_asset_filter(route): + req = route.request + if ( + req.resource_type in ("image", "font", "media", "stylesheet") + or "/_next/static/" in req.url + or "analytics" in req.url + ): + await route.abort() + return + await route.continue_() + + +async def grpc_verify_code(page, email, code): + inner = pb_str(1, email) + pb_str(2, code) + frame = b'\x00' + struct.pack('>I', len(inner)) + inner + fb64 = base64.b64encode(frame).decode() + s = await page.evaluate(f"(async()=>{{var fb=Uint8Array.from(atob('{fb64}'),c=>c.charCodeAt(0));var r=await fetch('{SITE_URL}/auth_mgmt.AuthManagement/VerifyEmailValidationCode',{{method:'POST',headers:{{'content-type':'application/grpc-web+proto','x-grpc-web':'1','x-user-agent':'connect-es/2.1.1'}},body:fb.buffer}});return r.headers.get('grpc-status')||'0';}})()") + if REGISTRATION_DIAGNOSTICS and s != '0': + debug_log('[C] verify rejected') + return s == '0' + + +def auth_cookie_snapshot(cookies): + """保留认证所需 Cookie 的原始作用域;不写入邮箱密码。""" + fields = ("name", "value", "domain", "path", "expires", "httpOnly", "secure", "sameSite") + return [ + {field: cookie[field] for field in fields if field in cookie} + for cookie in cookies + if cookie.get("name") in {"sso", "sso-rw"} + ] + + +async def server_action_register( + page, + email, + password, + code, + turnstile_token, + *, + include_session=False, +): + payload = json.dumps([{ + "emailValidationCode": code, + "createUserAndSessionRequest": { + "email": email, + "givenName": random.choice(["James","John","Robert","Michael","William","David","Richard","Joseph","Thomas","Charles"]), + "familyName": random.choice(["Smith","Johnson","Williams","Brown","Jones","Garcia","Miller","Davis","Rodriguez","Martinez"]), + "clearTextPassword": password, "tosAcceptedVersion": "$undefined" + }, + "turnstileToken": turnstile_token, "promptOnDuplicateEmail": True + }]) + pb64 = base64.b64encode(payload.encode()).decode() + if REGISTRATION_DIAGNOSTICS: + diagnostic_json = await page.evaluate(f"""(async()=>{{var r=await fetch('{SITE_URL}/sign-up',{{method:'POST',headers:{{'accept':'text/x-component','content-type':'text/plain;charset=UTF-8','next-router-state-tree':'{STATE_TREE}','next-action':'{ACTION_ID}'}},body:atob('{pb64}')}});return JSON.stringify({{status:r.status,retryAfter:r.headers.get('retry-after')||'',text:await r.text()}});}})()""") + diagnostic = json.loads(diagnostic_json) + result_text = diagnostic['text'] + else: + diagnostic = None + result_text = await page.evaluate(f"""(async()=>{{var r=await fetch('{SITE_URL}/sign-up',{{method:'POST',headers:{{'accept':'text/x-component','content-type':'text/plain;charset=UTF-8','next-router-state-tree':'{STATE_TREE}','next-action':'{ACTION_ID}'}},body:atob('{pb64}')}});return await r.text();}})()""") + # 注册响应里带一个 set-cookie 重定向 URL,必须访问它,x.ai 才会下发真正的 sso cookie(152 字符 JWT)。 + # 注意:直接解 q= 里的 JWT 取 config.token 是错的——那是 120 字符内部 blob,不是 sso 凭证。 + text = result_text.replace('\\/', '/') # RSC 里 / 被转义成 \/ + m = re.search(r'(https://[^" \s\\]+set-cookie\?q=[^:" \s\\]+)1:', text) + if not m: + m = re.search(r'(https://[^" \s\\]+set-cookie\?q=[A-Za-z0-9_.\-]+)', text) + if not m: + markers = _signup_response_markers(result_text) + if diagnostic is not None: + debug_log( + f"[C] signup no session http_status={diagnostic['status']} " + f"retry_after={diagnostic['retryAfter'] or '-'} response_bytes={len(result_text)} " + f"markers={markers}" + ) + if "rate_limited" in markers: + raise RegistrationRateLimited("signup_rate_limited") + return None + url = m.group(1) + if C_SET_COOKIE_VIA_REQUEST: + try: + await page.context.request.get(url, timeout=15000) + cookies = await page.context.cookies() + sso = next((c['value'] for c in cookies if c['name'] == 'sso'), None) + if sso: + if include_session: + return sso, auth_cookie_snapshot(cookies) + return sso + except asyncio.CancelledError: + raise + except Exception: + pass + + # 首方导航访问该 URL,浏览器正常落 sso cookie(跨域 fetch 会被 CORS/三方cookie 拦) + try: + await page.goto(url, timeout=15000, wait_until='domcontentloaded') + except asyncio.CancelledError: + raise + except Exception: + pass + cookies = await page.context.cookies() + sso = next((c['value'] for c in cookies if c['name'] == 'sso'), None) + if REGISTRATION_DIAGNOSTICS and not sso: + debug_log('[C] signup set-cookie completed without sso cookie') + if not sso: + return None + if include_session: + return sso, auth_cookie_snapshot(cookies) + return sso + +# solver 预热页面池:复用已停在 sign-up 的页面,省掉每次 page.goto 的重型 SPA 加载 +# SOLVER_REUSE=0 可关闭(用于 A/B 对比 goto 优化的增益) +SOLVER_REUSE = (os.environ.get("SOLVER_REUSE", "1").strip().lower() not in ("0", "false", "no")) +_solver_pool = [] +_solver_lock = asyncio.Lock() +MAX_SOLVER_REUSE = _env_int("MAX_SOLVER_REUSE", 25) +SOLVER_INITIAL_WAIT_MS = _env_int("SOLVER_INITIAL_WAIT_MS", 500) +SOLVER_POLL_INTERVAL_MS = _env_int("SOLVER_POLL_INTERVAL_MS", 500) +SOLVER_POLL_ATTEMPTS = _env_int("SOLVER_POLL_ATTEMPTS", 100) +SOLVER_HARD_TIMEOUT = max(10, _env_int("SOLVER_HARD_TIMEOUT", 90)) +SOLVER_CLEANUP_TIMEOUT = max(1, _env_int("SOLVER_CLEANUP_TIMEOUT", 5)) +SOLVER_FAST_CLICK = (os.environ.get("SOLVER_FAST_CLICK", "1").strip().lower() not in ("0", "false", "no")) +SOLVER_MOUSE_CLICK_RETRIES = _env_int("SOLVER_MOUSE_CLICK_RETRIES", 3) +SOLVER_MOUSE_CLICK_INTERVAL_MS = _env_int("SOLVER_MOUSE_CLICK_INTERVAL_MS", 600) +SOLVER_TIMELINE_TRACE = (os.environ.get("SOLVER_TIMELINE_TRACE", "0").strip().lower() in ("1", "true", "yes")) +SOLVER_TIMELINE_SAMPLE = _env_int("SOLVER_TIMELINE_SAMPLE", 8) +_solver_timeline_emitted = 0 +_solver_timeline_next_id = 0 + + +def _new_solver_timeline(*, enabled=None): + global _solver_timeline_next_id + enabled = SOLVER_TIMELINE_TRACE if enabled is None else enabled + if not enabled: + return None + _solver_timeline_next_id += 1 + return {"start": time.time(), "solve_id": _solver_timeline_next_id, "events": []} + + +def _trace_solver_event(timeline, event, **fields): + if timeline is None: + return + item = {"t": round(time.time() - timeline["start"], 4), "event": event} + if "solve_id" in timeline: + item["solve_id"] = timeline["solve_id"] + item.update(fields) + timeline["events"].append(item) + + +def _new_solver_poll_stats(): + return { + "poll_attempts": 0, + "first_token_attempt": None, + "poll_read_ms_total": 0.0, + "poll_read_ms_max": 0.0, + "poll_read_count": 0, + "poll_retry_click_count": 0, + } + + +def _finish_solver_poll_stats(stats): + if not stats: + return {} + read_count = stats.pop("poll_read_count", 0) + total = stats.pop("poll_read_ms_total", 0.0) + stats["poll_read_ms_avg"] = round(total / read_count, 1) if read_count else 0.0 + stats["poll_read_ms_max"] = round(stats.get("poll_read_ms_max", 0.0), 1) + return stats + + +async def _turnstile_frame_count(p): + try: + return await p.evaluate( + "() => document.querySelectorAll('iframe[src*=turnstile], iframe[src*=challenges.cloudflare.com]').length" + ) + except asyncio.CancelledError: + raise + except Exception: + return 0 + + +async def _turnstile_dom_snapshot(p): + try: + return await p.evaluate( + """() => { + const clip = (value, n = 96) => String(value || "").slice(0, n); + const rectInfo = (el) => { + if (!el) return null; + const r = el.getBoundingClientRect(); + return { + x: Math.round(r.left), + y: Math.round(r.top), + w: Math.round(r.width), + h: Math.round(r.height), + visible: r.width >= 1 && r.height >= 1 + }; + }; + const elemInfo = (el) => { + if (!el) return null; + return { + tag: clip(el.tagName), + id: clip(el.id, 64), + class: clip(el.className, 96), + is_iframe: el.tagName === "IFRAME" + }; + }; + const urlInfo = (src) => { + try { + const u = new URL(src || "", location.href); + return {host: clip(u.host, 96), path: clip(u.pathname, 96)}; + } catch (_) { + return {host: "", path: ""}; + } + }; + const widget = document.querySelector(".cf-turnstile"); + const wr = rectInfo(widget); + const center = wr ? { + x: Math.round(wr.x + wr.w / 2), + y: Math.round(wr.y + wr.h / 2) + } : null; + const centerEl = center ? document.elementFromPoint(center.x, center.y) : null; + const iframes = Array.from(document.querySelectorAll("iframe")); + const iframeSummaries = iframes.slice(0, 8).map((f) => { + const info = Object.assign(urlInfo(f.getAttribute("src") || ""), rectInfo(f) || {}); + info.in_widget = widget ? widget.contains(f) : false; + return info; + }); + const isTurnstileFrame = (f) => { + const src = f.getAttribute("src") || ""; + return src.includes("turnstile") || src.includes("challenges.cloudflare.com"); + }; + const response = document.querySelector('input[name="cf-turnstile-response"]'); + return { + __csp_solver_snapshot: true, + ready_state: document.readyState, + viewport: {w: window.innerWidth, h: window.innerHeight}, + widget: Object.assign({present: Boolean(widget)}, wr || {}), + click_center: center, + element_at_center: elemInfo(centerEl), + all_iframe_count: iframes.length, + turnstile_iframe_count: iframes.filter(isTurnstileFrame).length, + iframe_summaries: iframeSummaries, + turnstile_loaded: Boolean(window.turnstile), + response_input: { + present: Boolean(response), + token_len: response && response.value ? response.value.length : 0 + } + }; + }""" + ) + except asyncio.CancelledError: + raise + except Exception: + return {} + + +async def _turnstile_page_trace(p): + try: + return await p.evaluate( + "() => window.__cspTurnstileTrace ? Object.assign({}, window.__cspTurnstileTrace) : null" + ) + except asyncio.CancelledError: + raise + except Exception: + return None + + +async def _get_solver_page(browser): + if SOLVER_REUSE: + async with _solver_lock: + if _solver_pool: + item = _solver_pool.pop() + item["reused"] = True + item["goto_s"] = 0.0 + return item + p = await _prepare_browser_page(browser) + await p.set_viewport_size({"width": 800, "height": 600}) + goto_started = time.time() + await p.goto(f'{SITE_URL}/sign-up', timeout=20000) + await p.wait_for_timeout(1000) + return {"page": p, "n": 0, "reused": False, "goto_s": time.time() - goto_started} + +async def _put_solver_page(item, ok): + p = item["page"] + item["n"] += 1 + if SOLVER_REUSE and ok and item["n"] < MAX_SOLVER_REUSE: + try: # 清理本次注入痕迹,留待复用 + await asyncio.wait_for( + p.evaluate("document.querySelectorAll('.cf-turnstile').forEach(e=>e.remove());var i=document.querySelector('input[name=\"cf-turnstile-response\"]');if(i)i.remove();"), + timeout=SOLVER_CLEANUP_TIMEOUT, + ) + async with _solver_lock: + _solver_pool.append(item) + return + except asyncio.CancelledError: + await _close_solver_page(item) + raise + except Exception: + pass + await _close_solver_page(item) + + +async def _close_solver_page(item): + """Bound cleanup so a wedged renderer cannot trap the worker in finally.""" + page = item["page"] + try: + await asyncio.wait_for(page.close(), timeout=SOLVER_CLEANUP_TIMEOUT) + except asyncio.CancelledError: + # Solver cancellation is expected at the hard deadline. Give cleanup + # its own bounded task so the page is not returned to the reuse pool. + cleanup = asyncio.create_task(page.close()) + try: + await asyncio.wait_for(cleanup, timeout=SOLVER_CLEANUP_TIMEOUT) + except BaseException: + cleanup.cancel() + raise + except Exception: + pass + +async def _inject_turnstile_widget(p, *, timeline=False): + if not timeline: + await p.evaluate(f"""var d=document.createElement('div');d.className='cf-turnstile';d.setAttribute('data-sitekey','{SITE_KEY}');d.style.cssText='position:fixed;top:10px;left:10px;z-index:99999;background:white;padding:12px;border:2px solid red;border-radius:6px;width:300px;height:70px';document.body.appendChild(d);function __r(){{window.turnstile&&window.turnstile.render(d,{{sitekey:'{SITE_KEY}',callback:function(t){{var i=document.querySelector('input[name="cf-turnstile-response"]');if(!i){{i=document.createElement('input');i.type='hidden';i.name='cf-turnstile-response';document.body.appendChild(i);}}i.value=t;}}}})}}if(window.turnstile){{__r()}}else{{var s=document.createElement('script');s.src='https://challenges.cloudflare.com/turnstile/v0/api.js';s.onload=function(){{setTimeout(__r,1000)}};document.head.appendChild(s);}}""") + return + await p.evaluate(f"""var __trace=window.__cspTurnstileTrace={{created_at:performance.now(),script_inserted_at:null,script_loaded_at:null,render_called_at:null,render_returned_at:null,token_written_at:null,token_len:0,error:null}};var d=document.createElement('div');d.className='cf-turnstile';d.setAttribute('data-sitekey','{SITE_KEY}');d.style.cssText='position:fixed;top:10px;left:10px;z-index:99999;background:white;padding:12px;border:2px solid red;border-radius:6px;width:300px;height:70px';document.body.appendChild(d);function __r(){{try{{if(!window.turnstile)return;__trace.render_called_at=performance.now();var __ret=window.turnstile.render(d,{{sitekey:'{SITE_KEY}',callback:function(t){{var i=document.querySelector('input[name="cf-turnstile-response"]');if(!i){{i=document.createElement('input');i.type='hidden';i.name='cf-turnstile-response';document.body.appendChild(i);}}i.value=t;__trace.token_written_at=performance.now();__trace.token_len=t?t.length:0;}}}});__trace.render_returned_at=performance.now();__trace.render_return_type=typeof __ret;}}catch(e){{__trace.error=e&&e.name?e.name:String(e);}}}}if(window.turnstile){{__r()}}else{{var s=document.createElement('script');s.src='https://challenges.cloudflare.com/turnstile/v0/api.js';s.onload=function(){{__trace.script_loaded_at=performance.now();setTimeout(__r,1000)}};__trace.script_inserted_at=performance.now();document.head.appendChild(s);}}""") + + +async def _has_visible_turnstile_frame(p): + try: + return await p.evaluate( + """() => Array.from(document.querySelectorAll('iframe')).some((f) => { + const r = f.getBoundingClientRect(); + return r.width >= 20 && r.height >= 20; + })""" + ) + except Exception: + return False + + +async def _read_turnstile_token(p): + try: + return await p.evaluate('document.querySelector("input[name=\\"cf-turnstile-response\\"]")?.value||""') + except asyncio.CancelledError: + raise + except Exception: + return "" + + +async def _mouse_click_turnstile_center(p): + clicked, _trace = await _mouse_click_turnstile_center_trace(p) + return clicked + + +async def _mouse_click_turnstile_center_trace(p): + trace = {} + started = time.time() + box = await p.evaluate( + """() => { + const e = document.querySelector('.cf-turnstile'); + if (!e) return null; + const r = e.getBoundingClientRect(); + return {x: r.left + r.width / 2, y: r.top + r.height / 2}; + }""" + ) + trace["box_eval_ms"] = round((time.time() - started) * 1000, 1) + if not box: + return False, trace + x = float(box["x"]) + y = float(box["y"]) + trace["click_x"] = round(x, 1) + trace["click_y"] = round(y, 1) + started = time.time() + await p.mouse.move(max(0, x - 25), max(0, y - 8)) + trace["mouse_move1_ms"] = round((time.time() - started) * 1000, 1) + started = time.time() + await p.mouse.move(x, y, steps=8) + trace["mouse_move2_ms"] = round((time.time() - started) * 1000, 1) + started = time.time() + await p.mouse.down() + trace["mouse_down_ms"] = round((time.time() - started) * 1000, 1) + await asyncio.sleep(0.05) + started = time.time() + await p.mouse.up() + trace["mouse_up_ms"] = round((time.time() - started) * 1000, 1) + return True, trace + + +async def _repeat_mouse_click_turnstile(p, *, timeline=None): + retries = max(0, SOLVER_MOUSE_CLICK_RETRIES) + if retries <= 0: + return False + clicked = False + interval = max(50, SOLVER_MOUSE_CLICK_INTERVAL_MS) / 1000.0 + for i in range(retries): + token = await _read_turnstile_token(p) + dom = await _turnstile_dom_snapshot(p) if timeline is not None else None + iframe_count = dom.get("turnstile_iframe_count", 0) if dom else 0 + _trace_solver_event( + timeline, "click_before", attempt=i + 1, + token_len=len(token or ""), iframe_count=iframe_count, dom=dom + ) + if token and len(token) > 10: + return clicked + click_started = time.time() + click_error = None + try: + if timeline is not None: + attempt_clicked, click_trace = await _mouse_click_turnstile_center_trace(p) + else: + attempt_clicked = await _mouse_click_turnstile_center(p) + click_trace = {} + clicked = attempt_clicked or clicked + except asyncio.CancelledError: + raise + except Exception as exc: + click_error = type(exc).__name__ + attempt_clicked = False + click_trace = {} + click_call_ms = round((time.time() - click_started) * 1000, 1) + token = await _read_turnstile_token(p) if timeline is not None else "" + dom = await _turnstile_dom_snapshot(p) if timeline is not None else None + iframe_count = dom.get("turnstile_iframe_count", 0) if dom else 0 + _trace_solver_event( + timeline, "click_after", attempt=i + 1, clicked=clicked, + attempt_clicked=attempt_clicked, + token_len=len(token or ""), iframe_count=iframe_count, + click_call_ms=click_call_ms, click_error=click_error, + click_trace=click_trace, dom=dom + ) + if i != retries - 1: + await asyncio.sleep(interval) + return clicked + + +async def _click_turnstile_if_possible(p, *, fast=False, timeline=None): + if SOLVER_MOUSE_CLICK_RETRIES > 0: + return await _repeat_mouse_click_turnstile(p, timeline=timeline) + visible = await _has_visible_turnstile_frame(p) + _trace_solver_event(timeline, "visible_check", visible=visible) + if fast and not visible: + return visible + click_timeout = 500 if fast else 3000 + for sel in ["iframe[src*='challenges.cloudflare.com']","iframe[src*='turnstile']",".cf-turnstile iframe"]: + try: + fr = p.frame_locator(sel).first + await fr.locator("#checkbox, .checkbox, input[type=checkbox], body").first.click(timeout=click_timeout) + break + except asyncio.CancelledError: + raise + except Exception: + continue + return visible + + +async def _poll_turnstile_token(p, *, stats=None): + for i in range(SOLVER_POLL_ATTEMPTS): + await asyncio.sleep(max(50, SOLVER_POLL_INTERVAL_MS) / 1000) + if stats is not None: + stats["poll_attempts"] = i + 1 + read_started = time.time() + try: + t = await _read_turnstile_token(p) + if stats is not None: + read_ms = (time.time() - read_started) * 1000 + stats["poll_read_count"] += 1 + stats["poll_read_ms_total"] += read_ms + stats["poll_read_ms_max"] = max(stats["poll_read_ms_max"], read_ms) + if t and len(t) > 10: + if stats is not None and stats["first_token_attempt"] is None: + stats["first_token_attempt"] = i + 1 + return t + except asyncio.CancelledError: + raise + except Exception: + pass + retry_every = max(1, int(10000 / max(50, SOLVER_POLL_INTERVAL_MS))) + if i > 0 and i % retry_every == 0: + try: + await p.locator(".cf-turnstile").first.click(timeout=1000) + if stats is not None: + stats["poll_retry_click_count"] += 1 + except asyncio.CancelledError: + raise + except Exception: + pass + return None + + +async def _start_turnstile_challenge(browser, *, fast_click=False): + global _solver_timeline_emitted + item = await _get_solver_page(browser) + p = item["page"] + trace_timeline = SOLVER_TIMELINE_TRACE and _solver_timeline_emitted < SOLVER_TIMELINE_SAMPLE + if trace_timeline: + _solver_timeline_emitted += 1 + timeline = _new_solver_timeline(enabled=trace_timeline) + trace = { + "goto_s": item.get("goto_s", 0.0), + "reused": bool(item.get("reused", False)), + "reuse_count": item.get("n", 0), + "inject_s": 0.0, + "initial_s": 0.0, + "click_s": 0.0, + "wait_s": 0.0, + "visible_frame": False, + } + item["trace"] = trace + item["timeline"] = timeline + try: + stage_started = time.time() + _trace_solver_event(timeline, "inject_start") + await _inject_turnstile_widget(p, timeline=timeline is not None) + trace["inject_s"] = time.time() - stage_started + _trace_solver_event(timeline, "inject_done") + if timeline is not None: + _trace_solver_event( + timeline, "page_trace_after_inject", + page_trace=await _turnstile_page_trace(p) + ) + stage_started = time.time() + await p.wait_for_timeout(SOLVER_INITIAL_WAIT_MS) + trace["initial_s"] = time.time() - stage_started + _trace_solver_event(timeline, "initial_done") + if timeline is not None: + _trace_solver_event( + timeline, "page_trace_after_initial", + page_trace=await _turnstile_page_trace(p) + ) + stage_started = time.time() + trace["visible_frame"] = await _click_turnstile_if_possible( + p, fast=fast_click, timeline=timeline + ) + trace["click_s"] = time.time() - stage_started + _trace_solver_event(timeline, "click_stage_done", clicked=bool(trace["visible_frame"])) + if timeline is not None: + _trace_solver_event( + timeline, "page_trace_after_click", + page_trace=await _turnstile_page_trace(p) + ) + return item + except BaseException: + await _put_solver_page(item, False) + raise + + +async def _wait_turnstile_challenge(item): + ok = False + try: + wait_started = time.time() + timeline = item.get("timeline") + poll_stats = _new_solver_poll_stats() if timeline is not None else None + _trace_solver_event(timeline, "poll_start") + token = await _poll_turnstile_token(item["page"], stats=poll_stats) + item.get("trace", {})["wait_s"] = time.time() - wait_started + ok = token is not None + page_trace = await _turnstile_page_trace(item["page"]) if timeline is not None else None + _trace_solver_event( + timeline, "poll_done", ok=ok, token_len=len(token or ""), + page_trace=page_trace, **_finish_solver_poll_stats(poll_stats) + ) + return token + except asyncio.CancelledError: + raise + except Exception: + return None + finally: + timeline = item.get("timeline") + if timeline is not None: + debug_log("[solver_timeline] " + json.dumps(timeline["events"], separators=(",", ":"))) + await _put_solver_page(item, ok) + + +async def solve_one_turnstile(browser): + token, _trace = await solve_one_turnstile_with_trace(browser) + return token + + +async def solve_one_turnstile_with_trace(browser): + item = await _start_turnstile_challenge(browser, fast_click=SOLVER_FAST_CLICK) + token = await _wait_turnstile_challenge(item) + return token, item.get("trace", {}) + + +def _record_solver_trace(metrics, trace, total_seconds, token): + metrics.t_solve_count += 1 + metrics.t_solve_seconds += total_seconds + if token is None: + metrics.t_solve_failed += 1 + if not trace: + return + metrics.solver_goto_seconds += trace.get("goto_s", 0.0) + metrics.solver_inject_seconds += trace.get("inject_s", 0.0) + metrics.solver_initial_seconds += trace.get("initial_s", 0.0) + metrics.solver_click_seconds += trace.get("click_s", 0.0) + metrics.solver_wait_seconds += trace.get("wait_s", 0.0) + if trace.get("reused"): + metrics.solver_reused_count += 1 + if trace.get("visible_frame"): + metrics.solver_visible_frame_count += 1 + + +# ────────────────────────────────────────────── +# 邮箱服务:custom(自建 webhook) / tempmail(免费临时邮箱,多 provider fallback) +# ────────────────────────────────────────────── +# 免 key 的公共临时邮箱 provider(实测可用,互为 fallback 消灭单点): +# - mail.tm 同协议:mail.tm / mail.gw / duckmail.sbs +# - 独立 API:tempmail.lol +# handle 编码 provider,供 poll_code 分派;新增 provider 只要在这两处加一段即可。 +TEMPMAIL_BASES = ["https://api.mail.tm", "https://api.mail.gw", "https://api.duckmail.sbs"] + +def _extract_code(text): + """多层兜底提取验证码,抗邮件模板变化。""" + for pat in (r'>([A-Z0-9]{3}-[A-Z0-9]{3})<', r'>([A-Z0-9]{6})<', r'\b([A-Z0-9]{3}-?[A-Z0-9]{3})\b'): + m = re.search(pat, text) + if m: + return m.group(1).replace('-', '') + return None + +def _mailtm_create(base, password): + """mail.tm 同协议建箱;返回 (handle, email)。""" + d = req.get(f'{base}/domains', timeout=12).json() + d = d.get('hydra:member', d) if isinstance(d, dict) else d + doms = [x['domain'] for x in d if x.get('isActive', True) and not x.get('isPrivate', False)] + if not doms: + raise RuntimeError('no domain') + email = f'oc{secrets.token_hex(5)}@{doms[0]}' + req.post(f'{base}/accounts', json={'address': email, 'password': password}, timeout=12) + tok = req.post(f'{base}/token', json={'address': email, 'password': password}, timeout=12).json().get('token', '') + if not tok: + raise RuntimeError('no token') + return f'mt|{base}|{tok}', email + +def _lol_create(): + """tempmail.lol 建箱;返回 (handle, email)。""" + r = req.post('https://api.tempmail.lol/v2/inbox/create', timeout=12).json() + addr, tok = r.get('address', ''), r.get('token', '') + if not addr or not tok: + raise RuntimeError('lol create failed') + return f'lol|{tok}', addr + +def create_email(): + """custom 用自建域名(本地 webhook);tempmail 随机打散多个 provider,逐个 fallback。""" + if EMAIL_MODE == 'custom': + email = f'oc{secrets.token_hex(5)}@{EMAIL_DOMAIN}' + password = rand_str() + return email, email, password # 地址即用,验证码经 CF Worker POST 到本地 webhook + + password = rand_str() + # 优先 TempMail.lol(mail.tm 的 /token POST 在本机 fake-IP 隧道下 SSL 握手僵死、 + # req.post 的 timeout 不覆盖 _create 阶段 → 卡住永不 raise → try/except fallback + # 够不着「卡死」状态 → 第一 maker 把整条 P 焊死。故 lol 先跑,mail.tm 后置兜底)。 + makers = [_lol_create] + [(lambda b=b: _mailtm_create(b, password)) for b in TEMPMAIL_BASES] + for make in makers: + try: + handle, email = make() + return handle, email, password + except Exception: + continue + raise RuntimeError('所有临时邮箱 provider 均不可用') + +def _tempmail_fetch(handle): + """按 handle 前缀分派,取该邮箱当前邮件全文(subject+text+html);无则 None。""" + kind = handle.split('|', 1)[0] + if kind == 'lol': + tok = handle.split('|', 1)[1] + data = req.get(f'https://api.tempmail.lol/v2/inbox?token={tok}', timeout=10).json() + items = data.get('emails') or data.get('messages') or [] + if not items: + return None + return '\n'.join(f"{i.get('subject','')}\n{i.get('body','')}\n{i.get('html','')}" + for i in items if isinstance(i, dict)) + # mail.tm 同协议:handle = "mt|base|token" + _, base, tok = handle.split('|', 2) + hdr = {'Accept': 'application/json', 'Authorization': f'Bearer {tok}'} + data = req.get(f'{base}/messages', headers=hdr, timeout=10).json() + msgs = data if isinstance(data, list) else data.get('hydra:member', []) + if not msgs: + return None + mid = str(msgs[0].get('id') or '') + detail = req.get(f'{base}/messages/{mid}', headers=hdr, timeout=10).json() + parts = [str(detail.get(k, '')) for k in ['subject', 'intro', 'text', 'html']] + if isinstance(detail.get('html'), list): + parts.append('\n'.join(str(x) for x in detail['html'])) + return '\n'.join(parts) + +def poll_code(handle, max_wait=90): + """轮询验证码:custom 查本地 webhook /check;tempmail 按 provider 取信。""" + if EMAIL_MODE == 'custom': + for _ in range(max_wait): + time.sleep(1) + try: + resp = req.get(f'{LOCAL_EMAIL_API}/check/{handle}', timeout=5) + if resp.status_code == 200 and resp.json().get('code'): + return resp.json()['code'] + except Exception: + pass + return None + + for _ in range(max_wait): + time.sleep(1) + try: + text = _tempmail_fetch(handle) + if text: + code = _extract_code(text) + if code: + return code + except Exception: + pass + return None + + +async def _create_email_async(loop): + """在线程池中创建邮箱,避免阻塞 asyncio 事件循环。""" + return await loop.run_in_executor(POLL_EXECUTOR, create_email) + + +async def _poll_code_async(loop, handle): + """在线程池中轮询验证码。""" + return await loop.run_in_executor(POLL_EXECUTOR, poll_code, handle) + + +async def _acquire_many(sem, count): + """一次预留多个许可;取消或异常时回滚已获取许可。""" + acquired = 0 + try: + for _ in range(count): + await sem.acquire() + acquired += 1 + return acquired + except BaseException: + for _ in range(acquired): + sem.release() + raise + + +class _NoopAsyncSemaphore: + async def acquire(self): + return True + + def release(self): + return None + + +async def _send_q_request_batch(browser, physical_sem, p_send_sem, requests, metrics=None): + """使用一个页面发送一批 Q 请求。 + + 返回每个请求的 sent 状态。等待 Q 返回不在此函数内发生,因此这里释放 + Physical_Sem 后不会占用本地重资源。 + """ + p_send_acquired = False + physical_acquired = False + physical_wait_started = None + physical_hold_started = None + page = None + await p_send_sem.acquire() + p_send_acquired = True + try: + physical_wait_started = time.time() + await physical_sem.acquire() + physical_acquired = True + physical_hold_started = time.time() + if metrics is not None: + metrics.p_physical_count += 1 + metrics.p_physical_wait_seconds += physical_hold_started - physical_wait_started + page = await _prepare_browser_page(browser) + await page.set_viewport_size({"width": 800, "height": 600}) + try: + stage_started = time.time() + await _prepare_signup_page(page, redirect=True, timeout=30000) + if metrics is not None: + metrics.p_page_prepare_count += 1 + metrics.p_page_prepare_seconds += time.time() - stage_started + except asyncio.CancelledError: + raise + except Exception: + return [{**item, "sent": False} for item in requests] + + results = [] + send_started = time.time() + for item in requests: + sent = False + try: + sent = await grpc_create_code(page, item["email"]) + except asyncio.CancelledError: + raise + except Exception: + sent = False + results.append({**item, "sent": sent}) + if metrics is not None: + metrics.p_send_count += 1 + metrics.p_send_seconds += time.time() - send_started + return results + finally: + if page is not None: + try: + await page.close() + except Exception: + pass + if physical_acquired: + if metrics is not None and physical_hold_started is not None: + metrics.p_physical_hold_seconds += time.time() - physical_hold_started + physical_sem.release() + if p_send_acquired: + p_send_sem.release() + + +async def _poll_and_admit_q( + request, + inventory, + q_pending_sem, + q_slot_sem, + metrics, + *, + q_batch_lease=None, + admission_gate=None, +): + """等待单个 Q 返回并入库;每个请求独立释放 pending/inflight。""" + loop = asyncio.get_event_loop() + release_reservation = True + try: + try: + code = await asyncio.wait_for( + _poll_code_async(loop, request["handle"]), + timeout=P_REQUEST_TIMEOUT, + ) + except asyncio.CancelledError: + # poll_code 仍在运行时取消,底层轮询可能继续持有该请求;此处 + # 不能提前归还 pending/inflight,否则新请求会超量进入。 + release_reservation = False + raise + except asyncio.TimeoutError: + code = None + + if code is None: + metrics.q_discarded += 1 + return False + + metrics.q_returned += 1 + returned_at = time.time() + q_env = None + try: + q_env = await ResourceEnvelope.create_with_slot( + 'Q', + { + 'email': request["email"], + 'password': request["password"], + 'code': code, + }, + q_slot_sem, + expires_at=returned_at + Q_MAX_AGE, + ) + await inventory.put_q(q_env) + debug_log('[P] verification code admitted') + return True + except asyncio.CancelledError: + if q_env is not None and not q_env.released: + q_env.discard() + raise + except Exception: + if q_env is not None and not q_env.released: + q_env.discard() + metrics.q_discarded += 1 + return False + finally: + if release_reservation: + # poll 已终止(含超时/网络/解析异常),或 Q 已返回后任务在 + # 入库阶段取消:请求所有权已经回到本协程,必须归还许可。 + q_pending_sem.release() + if q_batch_lease is not None: + await q_batch_lease.release_one() + if admission_gate is not None: + await admission_gate.notify_changed() + + +def _observe_background_task(task): + """Consume detached task failures so background settlement is not silent.""" + try: + task.result() + except asyncio.CancelledError: + pass + except Exception as e: + debug_log(f'[P] background settle err: {sanitize_terminal_error(e)}') + + +# ────────────────────────────────────────────── +# CSP Worker +# ────────────────────────────────────────────── + +class _CHotPageLease: + def __init__(self, browser, metrics=None): + self.browser = browser + self.metrics = metrics + self.context = None + self.page = None + + async def __aenter__(self): + started = time.time() + try: + self.context, self.page = await _acquire_c_page(self.browser, self.metrics) + return self.page + finally: + if self.metrics is not None: + self.metrics.c_page_acquire_count += 1 + self.metrics.c_page_acquire_seconds += time.time() - started + + async def __aexit__(self, exc_type, exc, tb): + await _release_c_page(self.context, self.page, healthy=exc_type is None) + self.context = None + self.page = None + return False + + +_c_hot_page_pool = [] +_c_hot_page_lock = asyncio.Lock() +_c_hot_page_pool_size = derive_c_hot_page_pool_size(PHYSICAL_CAP or 1, C_WORKERS or 1) + + +async def _new_c_hot_page(browser): + context = await _prepare_browser_context(browser) + page = await context.new_page() + await page.set_viewport_size({"width": 800, "height": 600}) + await _prepare_signup_page(page, redirect=True, timeout=30000) + return context, page + + +async def _acquire_c_page(browser, metrics=None): + if C_HOT_PAGE_POOL: + async with _c_hot_page_lock: + if _c_hot_page_pool: + if metrics is not None: + metrics.c_hot_page_hits += 1 + return _c_hot_page_pool.pop() + if metrics is not None: + metrics.c_hot_page_misses += 1 + return await _new_c_hot_page(browser) + + page = await _prepare_browser_page(browser) + await page.set_viewport_size({"width": 800, "height": 600}) + await _prepare_signup_page(page, redirect=True, timeout=30000) + return None, page + + +async def _release_c_page(context, page, *, healthy): + if not C_HOT_PAGE_POOL or context is None: + try: + await page.close() + except Exception: + pass + return + + if healthy: + try: + if "/sign-up" not in (getattr(page, "url", "") or ""): + healthy = False + else: + await context.clear_cookies() + await page.evaluate( + "() => { try { localStorage.clear(); sessionStorage.clear(); } catch (e) {} }" + ) + async with _c_hot_page_lock: + if len(_c_hot_page_pool) < _c_hot_page_pool_size: + _c_hot_page_pool.append((context, page)) + return + except asyncio.CancelledError: + try: + await context.close() + except Exception: + pass + raise + except Exception: + pass + + try: + await context.close() + except Exception: + try: + await page.close() + except Exception: + pass + + +def _c_page_lease(browser, metrics=None): + return _CHotPageLease(browser, metrics) + + +async def _close_c_hot_page_pool(): + async with _c_hot_page_lock: + items = list(_c_hot_page_pool) + _c_hot_page_pool.clear() + for context, page in items: + try: + await context.close() + except Exception: + try: + await page.close() + except Exception: + pass + + +async def s_worker(wid, browser, inventory, physical_sem, t_slot_sem, metrics, admission_gate=None): + """S_Worker: 生成 T 并入库。""" + while not STOP.is_set(): + t_lease = None + try: + if admission_gate is not None: + t_lease = await admission_gate.acquire_t_production() + + physical_wait_started = time.time() + await physical_sem.acquire() + physical_hold_started = time.time() + metrics.s_physical_count += 1 + metrics.s_physical_wait_seconds += physical_hold_started - physical_wait_started + token = None + trace = {} + solve_started = time.time() + try: + try: + token, trace = await asyncio.wait_for( + solve_one_turnstile_with_trace(browser), + timeout=SOLVER_HARD_TIMEOUT, + ) + except asyncio.TimeoutError: + debug_log(f'[S] {wid} solver timeout after {SOLVER_HARD_TIMEOUT}s') + except asyncio.CancelledError: + raise + except Exception as exc: + debug_log(f'[S] {wid} solver error: {sanitize_terminal_error(exc)}') + finally: + solve_elapsed = time.time() - solve_started + _record_solver_trace(metrics, trace, solve_elapsed, token) + metrics.s_physical_hold_seconds += time.time() - physical_hold_started + physical_sem.release() + + if token is None: + metrics.t_discarded += 1 + if t_lease is not None: + await t_lease.release() + t_lease = None + await asyncio.sleep(0.5) + continue + + metrics.t_produced += 1 + now = time.time() + t_env = None + try: + t_env = await ResourceEnvelope.create_with_slot( + 'T', token, t_slot_sem, expires_at=now + T_MAX_AGE + ) + await inventory.put_t(t_env) + if admission_gate is not None: + await admission_gate.notify_changed() + except asyncio.CancelledError: + if t_env is not None and not t_env.released: + t_env.discard() + metrics.t_discarded += 1 + raise + except Exception: + if t_env is not None and not t_env.released: + t_env.discard() + metrics.t_discarded += 1 + except asyncio.CancelledError: + raise + except Exception as exc: + # A single browser/page failure must not terminate the permanent + # producer. The next iteration acquires a fresh solver page. + metrics.t_discarded += 1 + debug_log(f'[S] {wid} worker error: {sanitize_terminal_error(exc)}') + finally: + if t_lease is not None: + await t_lease.release() + await asyncio.sleep(0.2) + + +async def p_worker( + wid, + browser, + inventory, + physical_sem, + q_pending_sem, + q_slot_sem, + metrics, + admission_gate=None, + p_send_sem=None, + max_batch=1, +): + """P_Worker: 创建邮箱 + 发码 + 轮询 + 入库。""" + loop = asyncio.get_event_loop() + p_send_sem = p_send_sem or _NoopAsyncSemaphore() + while not STOP.is_set(): + q_lease = None + pending_owned = 0 + settle_tasks = [] + try: + if admission_gate is not None: + q_lease = await admission_gate.acquire_q_batch(max_batch=max_batch) + batch_count = q_lease.count + else: + batch_count = 1 + + pending_owned = await _acquire_many(q_pending_sem, batch_count) + + requests = [] + for _ in range(batch_count): + email_started = time.time() + try: + handle, email, password = await _create_email_async(loop) + metrics.p_email_create_count += 1 + metrics.p_email_create_seconds += time.time() - email_started + requests.append({"handle": handle, "email": email, "password": password}) + except asyncio.CancelledError: + raise + except Exception as e: + metrics.p_email_create_count += 1 + metrics.p_email_create_seconds += time.time() - email_started + debug_log(f'[P] {wid} create email err: {sanitize_terminal_error(e)}') + metrics.q_discarded += 1 + q_pending_sem.release() + pending_owned -= 1 + if q_lease is not None: + await q_lease.release_one() + + if not requests: + continue + + results = await _send_q_request_batch( + browser, physical_sem, p_send_sem, requests, metrics + ) + metrics.q_send_batches += 1 + metrics.q_send_batch_items += len(results) + + for item in results: + if not item["sent"]: + metrics.q_discarded += 1 + q_pending_sem.release() + pending_owned -= 1 + if q_lease is not None: + await q_lease.release_one() + continue + + metrics.q_sent += 1 + pending_owned -= 1 + task = asyncio.create_task( + _poll_and_admit_q( + item, + inventory, + q_pending_sem, + q_slot_sem, + metrics, + q_batch_lease=q_lease, + admission_gate=admission_gate, + ) + ) + task.add_done_callback(_observe_background_task) + settle_tasks.append(task) + + if settle_tasks: + await asyncio.gather(*(asyncio.shield(task) for task in settle_tasks)) + + except asyncio.CancelledError: + for _ in range(pending_owned): + q_pending_sem.release() + if q_lease is not None: + await q_lease.release_one() + raise + except Exception as e: + debug_log(f'[P] {wid} err: {sanitize_terminal_error(e)}') + metrics.q_discarded += 1 + for _ in range(pending_owned): + q_pending_sem.release() + if q_lease is not None: + await q_lease.release_one() + await asyncio.sleep(0.2) + + +def _pair_is_expired(pair, now=None): + """检查已 claim 的 T/Q 是否在等待期间失效。""" + now = time.time() if now is None else now + return any( + bool(check(now)) + for envelope in (pair.t, pair.q) + if (check := getattr(envelope, "is_expired", None)) is not None + ) + + +def _append_registration_line(path, line, mode=None, *, durable=False): + with open(path, "a") as stream: + stream.write(line) + if durable: + stream.flush() + os.fsync(stream.fileno()) + if mode is not None: + os.chmod(path, mode) + + +def _persist_registration(email, password, sso, session_cookies): + if session_cookies: + document = json.dumps( + {"email": email, "cookies": session_cookies}, + separators=(",", ":"), + ) + _append_registration_line( + "keys/auth-sessions.jsonl", + document + "\n", + mode=0o600, + durable=True, + ) + _append_registration_line("keys/accounts.txt", f"{email}:{password}:{sso}\n") + _append_registration_line("keys/grok.txt", sso + "\n") + + +async def _consume_pair( + browser, + physical_sem, + pair, + metrics, + task_id=None, + recovery_probe=None, +): + """执行一次 C 消费。返回 True 表示业务成功,False 表示消费失败。""" + global success_count + email = pair.q.value['email'] + password = pair.q.value['password'] + code = pair.q.value['code'] + token = pair.t.value + + # c_worker 在单次消费超时开始前完成冷却等待。保留这里的默认路径, + # 供直接调用者使用,同时避免把 60 秒冷却计入 60 秒消费超时。 + if recovery_probe is None: + recovery_probe = await REGISTRATION_RATE_LIMIT_CIRCUIT.wait() + + # 直接调用路径也可能在 circuit.wait() 中跨过资源有效期。尚未开始 + # 注册时只让出探针,不把本地过期误判成一次恢复探测失败。 + if _pair_is_expired(pair): + if recovery_probe: + REGISTRATION_RATE_LIMIT_CIRCUIT.release_probe(recovery_probe) + return False + + physical_acquired = False + physical_hold_started = None + recovery_completed = False + try: + physical_wait_started = time.time() + await physical_sem.acquire() + physical_acquired = True + physical_hold_started = time.time() + metrics.c_physical_count += 1 + metrics.c_physical_wait_seconds += physical_hold_started - physical_wait_started + t0 = time.time() + try: + async with _c_page_lease(browser, metrics) as page: + sso = None + session_cookies = [] + verify_started = time.time() + try: + verified = await grpc_verify_code(page, email, code) + finally: + metrics.c_verify_count += 1 + metrics.c_verify_seconds += time.time() - verify_started + # 限流可能由另一项并发任务在本任务验证码校验期间触发。 + # 非当前 probe 不得在 circuit 打开后开始新的注册提交;同时 + # 避免使用校验过程中刚过期的 T/Q。 + if ( + verified + and not _pair_is_expired(pair) + and REGISTRATION_RATE_LIMIT_CIRCUIT.can_submit(recovery_probe) + ): + register_started = time.time() + try: + registration = await server_action_register( + page, + email, + password, + code, + token, + include_session=True, + ) + if registration: + sso, session_cookies = registration + finally: + metrics.c_register_count += 1 + metrics.c_register_seconds += time.time() - register_started + except RegistrationRateLimited: + if REGISTRATION_RATE_LIMIT_CIRCUIT.trip(): + log( + format_user_registration_event( + "rate_limited", + wait_seconds=REGISTRATION_RATE_LIMIT_CIRCUIT.remaining_seconds(), + ) + ) + return False + + if sso: + elapsed = time.time() - t0 + async with file_lock: + _persist_registration(email, password, sso, session_cookies) + metrics.record_success() + success_count = metrics.success_count + count = metrics.success_count + log( + format_user_registration_event( + "success", + task_id=task_id, + count=count, + rate_per_minute=metrics.runtime_average_success_rate(), + ) + ) + if recovery_probe: + recovered_after = REGISTRATION_RATE_LIMIT_CIRCUIT.consume_recovery_seconds( + recovery_probe + ) + recovery_completed = True + if recovered_after is not None: + log( + format_user_registration_event( + "recovered", wait_seconds=round(recovered_after) + ) + ) + return True + + log(format_user_registration_event("failed", task_id=task_id)) + return False + finally: + if recovery_probe and not recovery_completed: + REGISTRATION_RATE_LIMIT_CIRCUIT.defer_probe(recovery_probe) + if physical_acquired: + metrics.c_physical_hold_seconds += time.time() - physical_hold_started + physical_sem.release() + + +async def c_worker(wid, browser, inventory, physical_sem, metrics, admission_gate=None): + """C_Worker: claim pair 并执行注册。""" + while not STOP.is_set(): + recovery_probe = False + task_id = None + try: + async with inventory.claim_pair() as pair: + if admission_gate is not None: + await admission_gate.notify_changed() + + # 必须先 claim,再通过冷却闸门。否则多个 worker 会在限流前 + # 通过闸门,随后于冷却期陆续拿到 pair 并漏闸。等待发生在 + # wait_for 之外,因此不计入单次 C 消费超时。 + recovery_probe = await REGISTRATION_RATE_LIMIT_CIRCUIT.wait() + if _pair_is_expired(pair): + if recovery_probe: + REGISTRATION_RATE_LIMIT_CIRCUIT.release_probe(recovery_probe) + recovery_probe = False + continue + + task_id = metrics.next_registration_task() + log( + format_user_registration_event( + "started", + task_id=task_id, + remaining=max(TARGET - metrics.success_count, 0) if TARGET else None, + ) + ) + try: + ok = await asyncio.wait_for( + _consume_pair( + browser, + physical_sem, + pair, + metrics, + task_id=task_id, + recovery_probe=recovery_probe, + ), + timeout=C_CONSUME_TIMEOUT, + ) + if ok: + metrics.pair_consumed_ok += 1 + else: + metrics.pair_consumed_fail += 1 + except asyncio.TimeoutError: + metrics.pair_consumed_fail += 1 + log(format_user_registration_event("failed", task_id=task_id)) + + except asyncio.CancelledError: + raise + except Exception as e: + if task_id is not None: + log(format_user_registration_event("failed", task_id=task_id)) + debug_log(f'[C] {wid} err: {sanitize_terminal_error(e)}') + metrics.pair_consumed_fail += 1 + finally: + if recovery_probe: + # 覆盖等待 pair、获取物理许可或其他边界处的取消/异常。 + # 已成功或已延期的 probe 在这里会是幂等 no-op。 + REGISTRATION_RATE_LIMIT_CIRCUIT.defer_probe(recovery_probe) + await asyncio.sleep(0.2) + + +# ────────────────────────────────────────────── +# 只读监控 +# ────────────────────────────────────────────── +async def monitor(inventory, sems, metrics, interval=8): + """定期输出系统状态。""" + while not STOP.is_set(): + await asyncio.sleep(interval) + if REGISTER_LOG_MODE == "debug": + log(metrics.snapshot(inventory, sems)) + if TARGET and metrics.success_count >= TARGET: + log(f'[*] 已达目标 {TARGET} 个,停止。'); STOP.set() + + +# ────────────────────────────────────────────── +# 主入口 +# ────────────────────────────────────────────── +async def main(): + global TARGET, _c_hot_page_pool_size + max_mem_arg = None + for i, arg in enumerate(sys.argv[1:], 1): + if arg == '--max-mem' and i + 1 < len(sys.argv): + max_mem_arg = sys.argv[i + 1] + elif arg == '--target' and i + 1 < len(sys.argv): + TARGET = int(sys.argv[i + 1]) + + resources = get_system_resources(max_mem_arg) + cpu = resources['cpu'] + capacity_profile = load_capacity_profile() + + # 自动派生容量 + physical_cap, s_workers, p_workers, c_workers = derive_capacity( + cpu, + resources['max_mem'], + profile_physical_cap=capacity_profile.get("physical_cap"), + ) + p_send_cap = P_SEND_CAP if P_SEND_CAP > 0 else 0 + admission_watermarks = derive_admission_watermarks(physical_cap) + _c_hot_page_pool_size = derive_c_hot_page_pool_size(physical_cap, c_workers) + + # 校验邮箱模式配置 + if EMAIL_MODE not in ('tempmail', 'custom'): + log("[!] 配置错误:EMAIL_MODE 应为 tempmail 或 custom"); return 2 + if EMAIL_MODE == 'custom' and not EMAIL_DOMAIN: + log("[!] 配置错误:custom 模式需在 .env 设置 EMAIL_DOMAIN"); return 2 + + debug_log("=" * 50) + debug_log(f" Grok Free Register (CSP Architecture)") + debug_log(f" CPU: {cpu} cores Memory: {resources['available_mem']}/{resources['total_mem']}MB") + debug_log(f" MaxMemForAuto: {resources['max_mem']}MB MemReserve: {MIN_FREE_MEM_MB}MB PhysicalMemBudget: {PHYSICAL_MEM_MB}MB") + if capacity_profile: + debug_log(f" CapacityProfile: {CAPACITY_PROFILE} physical_cap={capacity_profile.get('physical_cap')}") + debug_log(f" EmailMode: {EMAIL_MODE}") + debug_log(f" Physical_Sem={physical_cap} T_Slot={T_SLOT_CAP} Q_Slot={Q_SLOT_CAP} Q_Pending={Q_PENDING_CAP}") + debug_log( + f" Admission: T_LOW/HIGH={admission_watermarks['t_low']}/{admission_watermarks['t_high']} " + f"Q_LOW/HIGH={admission_watermarks['q_low']}/{admission_watermarks['q_high']}" + ) + debug_log(f" P_BatchMax={P_BATCH_MAX} P_Send_Sem={'disabled' if p_send_cap == 0 else p_send_cap}") + debug_log( + f" C_HotPagePool={'on' if C_HOT_PAGE_POOL else 'off'}" + f" size={_c_hot_page_pool_size if C_HOT_PAGE_POOL else 0}" + f" setCookieViaRequest={'on' if C_SET_COOKIE_VIA_REQUEST else 'off'}" + ) + debug_log(f" Workers: S={s_workers} P={p_workers} C={c_workers}") + debug_log( + f" Timeouts: Solver={SOLVER_HARD_TIMEOUT}s SolverCleanup={SOLVER_CLEANUP_TIMEOUT}s " + f"P_Request={P_REQUEST_TIMEOUT}s C_Consume={C_CONSUME_TIMEOUT}s" + ) + debug_log( + f" SolverMouseClick: retries={SOLVER_MOUSE_CLICK_RETRIES} " + f"interval={SOLVER_MOUSE_CLICK_INTERVAL_MS}ms" + ) + if TARGET: + debug_log(f" Target: {TARGET}") + reg_proxy = register_proxy_url() or "(direct)" + debug_log(f" RegisterProxy: {reg_proxy}") + debug_log("=" * 50) + + # 上次异常退出可能残留 CloakBrowser;先收尸再 launch,避免叠 3–4 份吃光内存 + reaped = await asyncio.to_thread(reap_orphan_register_browsers) + if reaped: + log(f"[*] 已回收 {reaped} 个残留注册浏览器进程") + + # CF prewarm via external FlareSolverr (:8191) — separate from REGISTER_PROXY egress + prewarm = await asyncio.to_thread(prewarm_clearance) + log(format_prewarm_log(prewarm)) + + await fetch_config() + + async with async_playwright() as pw: + browser = await pw.chromium.launch(**_launch_options()) + debug_log(f'[*] Browser launched (proxy={reg_proxy})') + + # CSP 组件 + metrics = Metrics() + inventory = Inventory(metrics=metrics) + admission_gate = AdmissionGate( + inventory, + t_low=admission_watermarks["t_low"], + t_high=admission_watermarks["t_high"], + q_low=admission_watermarks["q_low"], + q_high=admission_watermarks["q_high"], + ) + physical_sem = asyncio.Semaphore(physical_cap) + p_send_sem = asyncio.Semaphore(p_send_cap) if p_send_cap > 0 else None + t_slot_sem = asyncio.Semaphore(T_SLOT_CAP) + q_slot_sem = asyncio.Semaphore(Q_SLOT_CAP) + q_pending_sem = asyncio.Semaphore(Q_PENDING_CAP) + + sems = { + 'physical': physical_sem, + 't_slot': t_slot_sem, + 'q_slot': q_slot_sem, + 'q_pending': q_pending_sem, + 'admission': admission_gate, + } + if p_send_sem is not None: + sems['p_send'] = p_send_sem + + tasks = [] + + # S_Workers + for i in range(s_workers): + tasks.append(asyncio.create_task( + s_worker(i, browser, inventory, physical_sem, t_slot_sem, metrics, admission_gate) + )) + + # P_Workers + for i in range(p_workers): + tasks.append(asyncio.create_task( + p_worker( + i, + browser, + inventory, + physical_sem, + q_pending_sem, + q_slot_sem, + metrics, + admission_gate, + p_send_sem, + P_BATCH_MAX, + ) + )) + + # C_Workers + for i in range(c_workers): + tasks.append(asyncio.create_task( + c_worker(i, browser, inventory, physical_sem, metrics, admission_gate) + )) + + # Monitor + tasks.append(asyncio.create_task(monitor(inventory, sems, metrics))) + + debug_log(f'[*] CSP up: S={s_workers} P={p_workers} C={c_workers} workers') + log( + format_user_registration_event( + "service_started", + remaining=max(TARGET - metrics.success_count, 0) if TARGET else None, + ) + ) + + try: + await asyncio.gather(*tasks) + except (KeyboardInterrupt, asyncio.CancelledError): + pass + finally: + for t in tasks: + t.cancel() + with suppress(Exception): + await asyncio.wait_for(asyncio.gather(*tasks, return_exceptions=True), timeout=3) + await _close_c_hot_page_pool() + await _close_browser_hard(browser, timeout=5.0) + # 双保险:close 后仍扫一遍孤儿 + leftover = await asyncio.to_thread(reap_orphan_register_browsers) + if leftover: + log(f"[*] 退出时补杀 {leftover} 个残留浏览器") + log(format_user_registration_event("stopped", count=metrics.success_count)) + return 0 + +if __name__ == "__main__": + try: + REGISTER_LOG_MODE = resolve_register_log_mode(sys.argv[1:]) + except ValueError: + log("[!] 配置错误:REGISTER_LOG_MODE 应为 user 或 debug") + raise SystemExit(2) + try: + exit_code = asyncio.run(main()) + except ValueError as exc: + log(f"[!] 配置错误:{sanitize_terminal_error(exc)}") + exit_code = 2 + except KeyboardInterrupt: + # asyncio.run 外层再兜一层;收尸防止 kill -2 叠进程 + try: + n = reap_orphan_register_browsers() + if n: + print(f"[*] 中断收尸 {n} 个浏览器", flush=True) + except Exception: + pass + exit_code = 130 + raise SystemExit(exit_code) diff --git a/keys/acpa_watchdog.py b/keys/acpa_watchdog.py new file mode 100755 index 0000000..7d6f04e --- /dev/null +++ b/keys/acpa_watchdog.py @@ -0,0 +1,944 @@ +#!/usr/bin/env python3 +""" +acpa_watchdog.py — 自动 CPA 凭据整备驻守进程 + +职责: + 1. 盯 xai_enroller (auth-service.sh) 的出货目录: + $XAI_ENROLLER_LOCAL_AUTH_DIR/authenticated/xai-*.json + 默认 ~/Downloads/grok-free-register-auth/authenticated/ + xai_enroller 产出的 CPA 文件 base_url=api.x.ai/v1 (付费口, 免费号必 402) + + headers=None (缺 grok-cli 身份头) + email=None。 + 2. 每个新号原样不动地「整备」一份到 PROJECT/keys/cpa_ready/xai-*.json: + - base_url → https://cli-chat-proxy.grok.com/v1 (免费 Grok 4.5 口) + - headers → grok-cli 身份头(x-grok-client-version/identifier 等) + - email → 从 id_token /access_token JWT 里挖(便于辨识, 不影响通信用) + - 权限 0600 (合规: 含 refresh_token) + 3. 探活: 打 cli-chat-proxy /v1/responses model=grok-4.5 + 请求对齐 new-api grok-cli 完整指纹(session 族头 + input 数组 + max_output_tokens 极小)。 + 首探 ≤MAX_PROBES(默认 1,省 free 2M 额度);200 → alive。 + 新号落盘后 ACPA_PROBE_WARMUP_SEC(默认 3s) 再探,避免 mint 后瞬时 permission-denied。 + 单次探针内 403 当场短重试 ACPA_403_IMMEDIATE_RETRIES×SLEEP(默认 2×4s)。 + 仍 403 → chat_denied:不丢,RETEST_AFTER_SEC(默认 180s) 后回测 1 次; + 回测活 → alive;仍 403 → chat_dead 丢。 + 429 free-usage-exhausted / rate_limited / 401 → 立刻丢。 + alive 永不复探。 + 4. 状态: keys/cpa_ready/_state.tsv + name\tsub\temail\tstatus\tts\tprobes + 每轮 prune:json 在 cpa_ready/_discarded 都不存在的行直接从 TSV 删掉 + (删库跑路后不会残留幽灵 alive;盘上有 json 但无 state 会重新进探活)。 + +只读 AUTH_DIR, 只写 keys/cpa_ready/。源 CPA 文件不动; 不重铸 device flow; +不做 HTTP 灌库 (入库步骤留给灌库脚本, 此处只把号整备到可直接喂 cliproxy / new-api)。 + +CLI: + python3 acpa_watchdog.py # 常驻轮询 (默认 2s), Ctrl-C 退 + python3 acpa_watchdog.py --once # 单次扫一轮现存号就退 (CLI 自检) + python3 acpa_watchdog.py --interval 4 +""" +from __future__ import annotations + +import argparse +import base64 +import copy +import json +import os +import signal +import sys +import time +from pathlib import Path + +# --------------------------------------------------------------------------- +# paths — portable: follow project root of this script; AUTH overridable via env +# --------------------------------------------------------------------------- +PROJECT_ROOT = Path(__file__).resolve().parent.parent # .../grok-free-register +# Align with xai_enroller.service DEFAULT_LOCAL_AUTH_DIR +# (~/Downloads/grok-free-register-auth/authenticated by default) +_AUTH_BASE = Path( + os.environ.get( + "XAI_ENROLLER_LOCAL_AUTH_DIR", + os.environ.get( + "XAI_AUTH_SERVICE_LOCAL_DIR", + str(Path.home() / "Downloads" / "grok-free-register-auth"), + ), + ) +).expanduser() +AUTH_DIR = _AUTH_BASE / "authenticated" +OUT_DIR = PROJECT_ROOT / "keys" / "cpa_ready" +DISCARD_DIR = OUT_DIR / "_discarded" +STATE_FILE = OUT_DIR / "_state.tsv" + +# grok-cli 静态身份头(对齐 new-api common.GrokCLI* / SetupRequestHeader) +# session/req/agent 等 per-request 头只在探针里生成,不写进落盘 json。 +CLIPROXY_BASE_URL = "https://cli-chat-proxy.grok.com/v1" +CLIPROXY_HEADERS = { + "x-grok-client-version": "0.2.93", + "x-xai-token-auth": "xai-grok-cli", + "X-XAI-Token-Auth": "xai-grok-cli", + "x-authenticateresponse": "authenticate-response", + "x-grok-client-identifier": "grok-shell", + "x-compaction-at": "400000", + "User-Agent": "grok-shell/0.2.93 (linux; x86_64)", +} + +RUN = True + + +def _sig(signum, _frame): + global RUN + RUN = False + + +signal.signal(signal.SIGINT, _sig) +signal.signal(signal.SIGTERM, _sig) + + +# --------------------------------------------------------------------------- +# helpers +# --------------------------------------------------------------------------- +def log(msg: str) -> None: + print(f"[acpa {time.strftime('%H:%M:%S')}] {msg}", flush=True) + + +def b64url_decode(seg: str) -> bytes: + seg += "=" * (-len(seg) % 4) + return base64.urlsafe_b64decode(seg) + + +def jwt_payload(token: str) -> dict: + if not token or token.count(".") != 2: + return {} + try: + return json.loads(b64url_decode(token.split(".")[1])) + except Exception: + return {} + + +def load_source(path: Path) -> dict | None: + try: + return json.loads(path.read_text(encoding="utf-8")) + except Exception as exc: + log(f" ✗ 读 {path.name} 失败: {exc}") + return None + + +def ensure_dir(p: Path) -> None: + p.mkdir(parents=True, exist_ok=True) + try: + os.chmod(p, 0o700) + except OSError: + pass + + +def fallback_email(src: dict) -> str: + """xai_enroller 产的 file email 字段是 None。从 JWT 里挖一个标识。""" + for key in ("id_token", "access_token"): + pl = jwt_payload(src.get(key) or "") + for f in ("email", "preferred_username", "sub"): + v = pl.get(f) + if v: + return str(v) + return "" + + +def finalize(src: dict) -> tuple[str, dict]: + """把 xai_enroller 产的 raw CPA 整备成 cli-chat-proxy 可用格式。返回 (sub, dst)。""" + dst = copy.deepcopy(src) + dst["base_url"] = CLIPROXY_BASE_URL + dst["headers"] = dict(CLIPROXY_HEADERS) + # token_endpoint 不变 (auth.x.ai/oauth2/token, refresh 用) + dst.setdefault("auth_kind", "oauth") + dst.setdefault("type", "xai") + if not dst.get("email"): + dst["email"] = fallback_email(src) + sub = dst.get("sub") or jwt_payload(dst.get("access_token", "")).get("sub", "") + dst["sub"] = sub + return sub, dst + + +def write_out(name: str, entry: dict) -> Path: + ensure_dir(OUT_DIR) + target = OUT_DIR / name + tmp = target.with_suffix(target.suffix + ".tmp") + tmp.write_text(json.dumps(entry, separators=(",", ":"), ensure_ascii=False), + encoding="utf-8") + os.replace(tmp, target) + os.chmod(target, 0o600) + return target + + +# --------------------------------------------------------------------------- +# probe liveness — 使用项目自带 .venv(已装 curl_cffi) +# 完整 grok-cli 指纹;默认每号只首探 1 次(free ~2M,禁止连 ping) +# --------------------------------------------------------------------------- +_VENV_PY = PROJECT_ROOT / ".venv" / "bin" / "python" +_PROBE_SRC = r''' +import json, sys, uuid +try: + from curl_cffi import requests +except Exception as e: + print("NO_CURL_CFFI", str(e)); sys.exit(3) + +fp = sys.argv[1] +d = json.load(open(fp)) +at = d.get("access_token") or "" +if not at: + print(json.dumps({"status": "token_dead", "http": 401, "snippet": "no access_token"})) + sys.exit(0) + +# 对齐 new-api-fusion relay/channel/grok_cli SetupRequestHeader + common.GrokCLI* +sid = str(uuid.uuid4()) +rid = str(uuid.uuid4()) +hdrs = { + "Authorization": "Bearer " + at, + "Content-Type": "application/json", + "Accept": "application/json", + "User-Agent": "grok-shell/0.2.93 (linux; x86_64)", + "x-xai-token-auth": "xai-grok-cli", + "X-XAI-Token-Auth": "xai-grok-cli", + "x-grok-client-identifier": "grok-shell", + "x-grok-client-version": "0.2.93", + "x-authenticateresponse": "authenticate-response", + "x-compaction-at": "400000", + "Connection": "Keep-Alive", + "x-grok-session-id": sid, + "x-grok-conv-id": sid, + "x-grok-req-id": rid, + "x-grok-turn-idx": "1", + "x-grok-agent-id": "agent-" + rid[:8], + "x-grok-model-override": "grok-4.5", +} +email = d.get("email") or "" +sub = d.get("sub") or "" +if email: + hdrs["x-email"] = str(email) +if sub: + hdrs["x-userid"] = str(sub) + +base = (d.get("base_url") or "https://cli-chat-proxy.grok.com/v1").rstrip("/") +# 兼容 base 已带 /v1 或只有 host +if base.endswith("/responses"): + url = base +else: + url = base + "/responses" + +# input 必须是 items 数组;裸字符串会被上游拒 / 误 403 +# max_output_tokens 压到极小,free 号额度只花探针这一枪 +body = { + "model": "grok-4.5", + "store": False, + "stream": False, + "max_output_tokens": 16, + "input": [ + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "ok"}], + } + ], +} +try: + r = requests.post(url, json=body, headers=hdrs, impersonate="chrome", timeout=45) + txt = r.text or "" + code = r.status_code +except Exception as e: + print("EXC", str(e)[:200]); sys.exit(2) + +status = "unknown" +low = txt.lower() +if code == 200 and ("output" in txt or "completed" in low or '"status"' in low): + # 200 即视为通路可用;不要求模型真回 "ok"(省解码/推理) + status = "alive" +elif ( + code == 429 + or "free-usage-exhausted" in low + or "used all the include" in low + or "rate limit" in low + or "rate_limit" in low +): + status = "free_exhausted" if ("free-usage" in low or "used all" in low) else "rate_limited" +elif code == 402 or "personal-team-blocked" in txt or "spending-limit" in txt: + status = "no_quota_paid_end" +elif code == 403 or "permission-denied" in txt or "chat endpoint is denied" in low: + status = "chat_denied" +elif code == 401: + status = "token_dead" +elif code == 422: + # body 形态问题:标 unknown 便于发现探针回归,不直接当号死 + status = "unknown" +elif "forbidden" in low: + status = "chat_denied" +print(json.dumps({"status": status, "http": code, "snippet": txt[:200]})) +sys.exit(0) +''' + +# 省额度:默认首探 1 次;403 当场短重试后再标 soft,默认 3 分钟后延迟回测 1 次。 +# alive 永不复探。429 free-usage-exhausted / token_dead 等真坏号立刻丢。 +MAX_PROBES = int( + os.environ.get( + "ACPA_MAX_PROBES", + os.environ.get("ACPA_PROBE_RETRIES", "1"), + ) +) +PROBE_RETRY_SLEEP = float(os.environ.get("ACPA_PROBE_RETRY_SLEEP", "1.0")) +# 新号 finalize 后稍等再首探:上游常对刚 mint 的 access 回瞬时 permission-denied +PROBE_WARMUP_SEC = float(os.environ.get("ACPA_PROBE_WARMUP_SEC", "3.0")) +# 单次 probe_one 内 403 当场短重试(不算额外 MAX_PROBES 预算) +IMMEDIATE_403_RETRIES = int(os.environ.get("ACPA_403_IMMEDIATE_RETRIES", "2")) +IMMEDIATE_403_SLEEP = float(os.environ.get("ACPA_403_IMMEDIATE_SLEEP", "4.0")) +RETEST_AFTER_SEC = int(os.environ.get("ACPA_RETEST_AFTER_SEC", "180")) # 3 分钟 +RETEST_MAX = int(os.environ.get("ACPA_RETEST_MAX", "1")) # 403 延迟回测次数 + +# 终态:不再探(alive 保留;下列坏号 discard) +TERMINAL_OK = frozenset({"alive"}) +TERMINAL_BAD = frozenset({ + "free_exhausted", # 429 免费额度耗尽 → 丢 + "rate_limited", # 其它 429 → 丢 + "token_dead", # 401 + "no_quota_paid_end", + "chat_dead", # 5 分钟回测后仍 403 → 真坏,丢 + "discarded", +}) +TERMINAL_STATUSES = TERMINAL_OK | TERMINAL_BAD + +# 仅这些立刻/确认后丢出 cpa_ready(403 首探不在此列) +DISCARD_STATUSES = frozenset({ + "free_exhausted", + "rate_limited", + "token_dead", + "no_quota_paid_end", + "chat_dead", +}) + +# 等 5 分钟回测的软状态(保留 json,不进 acc,不丢) +SOFT_RETEST_STATUSES = frozenset({ + "chat_denied", + "forbidden_domain_ban", # 历史误标,按 403 软状态回测 + "probe_cap", # 历史:首探 2 次 403 被钉死,捞回可回测 +}) + + +def probe_one(path: Path) -> dict: + """用子进程跑探活(隔离 curl_cffi 环境)。返回 {status, http, snippet}.""" + import subprocess + cp = subprocess.run( + [str(_VENV_PY), "-c", _PROBE_SRC, str(path)], + capture_output=True, text=True, timeout=60, + ) + out = cp.stdout.strip() + if out.startswith("alive") or out.lstrip().startswith("{"): + try: + idx = out.find("{") + if idx >= 0: + return json.loads(out[idx:]) + except Exception: + pass + return {"status": "probe_err", "http": -1, "snippet": (out or cp.stderr[:200])[:200]} + + +def probe_one_soft(path: Path) -> dict: + """ + 一次「逻辑探针」:底层请求 + 403/瞬时错误当场短重试。 + 仍计为 probe_budgeted 的 1 次 used(不因短重试烧额外额度记账)。 + """ + last = {"status": "probe_err", "http": -1, "snippet": ""} + attempts = 1 + max(0, IMMEDIATE_403_RETRIES) + for i in range(attempts): + last = probe_one(path) + st = last.get("status") + if st in ( + "alive", "free_exhausted", "rate_limited", + "no_quota_paid_end", "token_dead", + ): + return last + if st in ("chat_denied", "forbidden_domain_ban", "probe_err", "unknown"): + if i + 1 < attempts: + time.sleep(IMMEDIATE_403_SLEEP) + continue + break + if last.get("status") == "forbidden_domain_ban": + last = { + "status": "chat_denied", + "http": last.get("http", 403), + "snippet": last.get("snippet", "")[:200], + } + return last + + +def probe_budgeted(path: Path, budget: int) -> tuple[dict, int]: + """ + 在 budget 次内探活。返回 (最后结果, 实际探活次数)。 + alive / 额度终态立刻停;403 可在预算内连探;最终 403 保持 chat_denied(不改 probe_cap)。 + """ + budget = max(0, int(budget)) + if budget <= 0: + return {"status": "chat_denied", "http": -1, "snippet": "no probe budget"}, 0 + last = {"status": "probe_err", "http": -1, "snippet": ""} + used = 0 + for i in range(budget): + last = probe_one_soft(path) + used += 1 + st = last.get("status") + if st in ( + "alive", "free_exhausted", "rate_limited", + "no_quota_paid_end", "token_dead", + ): + return last, used + if i + 1 < budget and st in ( + "chat_denied", "forbidden_domain_ban", "probe_err", "unknown", + ): + time.sleep(PROBE_RETRY_SLEEP) + continue + break + # 首探/回测结束仍是 403 → 保持 chat_denied,交给延迟回测或 chat_dead + if last.get("status") in ("forbidden_domain_ban",): + last = { + "status": "chat_denied", + "http": last.get("http", 403), + "snippet": last.get("snippet", "")[:200], + } + return last, used + + +# --------------------------------------------------------------------------- +# state — name \t sub \t email \t status \t ts [\t probes] +# --------------------------------------------------------------------------- +def load_state() -> dict: + """{name: {sub, email, status, ts, probes}}""" + if not STATE_FILE.exists(): + return {} + st = {} + for ln in STATE_FILE.read_text(encoding="utf-8").splitlines(): + parts = ln.split("\t") + if len(parts) >= 5: + name, sub, email, status, ts = parts[:5] + probes = 0 + if len(parts) >= 6: + try: + probes = int(parts[5] or 0) + except ValueError: + probes = 0 + # 历史行没有 probes:已有终态按已探满处理,避免重启后重烧额度 + if probes <= 0 and status in TERMINAL_STATUSES | { + "chat_denied", "forbidden_domain_ban", "unknown", "probe_err", + }: + probes = MAX_PROBES if status != "alive" else 1 + if probes <= 0 and status == "alive": + probes = 1 + st[name] = { + "sub": sub, "email": email, "status": status, + "ts": ts, "probes": probes, + } + return st + + +def write_state(st: dict) -> None: + ensure_dir(OUT_DIR) + lines = [] + for name, v in sorted(st.items()): + lines.append("\t".join([ + name, v.get("sub", ""), v.get("email", ""), + v.get("status", ""), str(v.get("ts", "")), + str(int(v.get("probes") or 0)), + ])) + tmp = STATE_FILE.with_suffix(".tsv.tmp") + tmp.write_text("\n".join(lines) + ("\n" if lines else ""), encoding="utf-8") + os.replace(tmp, STATE_FILE) + os.chmod(STATE_FILE, 0o600) + + +def json_exists_for(name: str) -> bool: + """cpa_ready 主目录或 _discarded 里是否还有该号 json。""" + return (OUT_DIR / name).exists() or (DISCARD_DIR / name).exists() + + +def prune_orphan_state(st: dict) -> int: + """ + 删库跑路 / 手工清 json 后:TSV 里既无 cpa_ready 也无 _discarded 文件的行直接 drop。 + 避免幽灵 alive 把 sync/账本钉死。返回删除行数。 + """ + if not st: + return 0 + drop = [name for name in list(st.keys()) if not json_exists_for(name)] + if not drop: + return 0 + for name in drop: + del st[name] + write_state(st) + log(f" 🧹 prune orphan state {len(drop)} 行(盘上无 json)") + return len(drop) + + +def should_discard(status: str, probes: int = 0) -> bool: + """仅真坏号丢:429 额度耗尽 / 401 / 回测后仍 403(chat_dead)。首探 403 不丢。""" + return status in DISCARD_STATUSES + + +def discard_bad(name: str, st: dict, *, reason: str = "") -> bool: + """ + 坏号移出 keys/cpa_ready/ → keys/cpa_ready/_discarded/ + state 标记 discarded,sync_acc 只认 alive 故 acc.md 自动剔除。 + 返回是否执行了丢弃。 + """ + entry = st.get(name) or {} + status = entry.get("status", reason or "bad") + probes = _probes(entry) if entry else MAX_PROBES + if status == "discarded": + # 已丢过:确保 json 不在主目录 + p = OUT_DIR / name + if p.exists(): + ensure_dir(DISCARD_DIR) + dest = DISCARD_DIR / name + try: + os.replace(p, dest) + except OSError: + try: + p.unlink() + except OSError: + pass + return False + if not should_discard(status, probes) and not reason: + return False + + ensure_dir(DISCARD_DIR) + src = OUT_DIR / name + moved = False + if src.exists(): + dest = DISCARD_DIR / name + try: + if dest.exists(): + dest.unlink() + os.replace(src, dest) + os.chmod(dest, 0o600) + moved = True + except OSError as exc: + log(f" ✗ 丢弃移动失败 {name}: {exc}") + try: + src.unlink() + moved = True + except OSError: + pass + + st[name] = { + "sub": entry.get("sub", ""), + "email": entry.get("email", ""), + "status": "discarded", + "ts": int(time.time()), + "probes": probes if probes else MAX_PROBES, + } + # 旁路账本:便于扫为何丢 + try: + ledger = DISCARD_DIR / "_discarded.tsv" + with ledger.open("a", encoding="utf-8") as fh: + fh.write( + f"{int(time.time())}\t{name}\t{entry.get('email','')}\t" + f"{status}\t{reason or status}\n" + ) + os.chmod(ledger, 0o600) + except OSError: + pass + log( + f" 🗑 discard {name} [{status}]" + f"{' → _discarded/' if moved else ' (no json)'}" + f" email={entry.get('email','')}" + ) + return True + + +def sweep_discard(st: dict) -> int: + """扫 state + cpa_ready,丢掉所有真坏号 json。返回丢弃个数。""" + n = 0 + for name in list(st.keys()): + if not RUN: + break + v = st[name] + if v.get("status") == "discarded": + discard_bad(name, st) # 清残留 json + continue + # 软 403 已用完延迟回测次数 → 升格 chat_dead 再丢 + if v.get("status") in SOFT_RETEST_STATUSES: + if _retests_done(v) >= RETEST_MAX and _age_sec(v) >= RETEST_AFTER_SEC: + v["status"] = "chat_dead" + st[name] = v + if should_discard(v.get("status", ""), _probes(v)): + if discard_bad(name, st): + n += 1 + write_state(st) + return n + + +def rescue_soft_discards(st: dict) -> int: + """ + 把误丢的 403/probe_cap 从 _discarded 捞回主目录,标 chat_denied, + 立刻可做一次延迟回测(ts 回拨到已到期)。 + 账本原因为 soft,或 state 已 discarded 但 json 仍在 _discarded 且无硬坏标记。 + """ + ensure_dir(DISCARD_DIR) + ensure_dir(OUT_DIR) + soft_reasons = SOFT_RETEST_STATUSES | {"probe_cap", "chat_denied", "forbidden_domain_ban"} + soft_names: set[str] = set() + ledger = DISCARD_DIR / "_discarded.tsv" + if ledger.exists(): + for ln in ledger.read_text(encoding="utf-8").splitlines(): + parts = ln.split("\t") + if len(parts) >= 5: + name, prev_st, reason = parts[1], parts[3], parts[4] + if prev_st in soft_reasons or reason in soft_reasons: + soft_names.add(name) + elif len(parts) >= 4: + name, prev_st = parts[1], parts[3] + if prev_st in soft_reasons: + soft_names.add(name) + + n = 0 + for p in sorted(DISCARD_DIR.glob("xai-*.json")): + if not RUN: + break + name = p.name + # 账本标 soft,或 state 缺/已 discarded 时默认按 soft 捞(旧误丢) + entry = st.get(name) or {} + prev = entry.get("status", "discarded") + if name not in soft_names and prev not in ("discarded", "") and prev not in soft_reasons: + continue + if name not in soft_names and prev == "discarded": + # 无账本线索:仍捞(用户确认 403 常误杀);真 429 已不在 soft 集合时靠账本 + soft_names.add(name) + if name not in soft_names: + continue + dest = OUT_DIR / name + try: + if dest.exists(): + dest.unlink() + os.replace(p, dest) + os.chmod(dest, 0o600) + except OSError as exc: + log(f" ✗ rescue 移动失败 {name}: {exc}") + continue + # ts 回拨:立刻进入可回测窗口(不额外空等 RETEST_AFTER_SEC) + st[name] = { + "sub": entry.get("sub", ""), + "email": entry.get("email", ""), + "status": "chat_denied", + "ts": int(time.time()) - RETEST_AFTER_SEC, + "probes": MAX_PROBES, # 留给 1 次回测额度(与当前 MAX_PROBES 对齐) + } + n += 1 + log(f" ↩ rescue {name} → chat_denied (retest due) probes={MAX_PROBES}") + if n: + write_state(st) + return n + + +# --------------------------------------------------------------------------- +# main loop +# --------------------------------------------------------------------------- +def _probes(entry: dict) -> int: + try: + return int(entry.get("probes") or 0) + except (TypeError, ValueError): + return 0 + + +def _age_sec(entry: dict) -> int: + try: + return max(0, int(time.time()) - int(entry.get("ts") or 0)) + except (TypeError, ValueError): + return RETEST_AFTER_SEC + + +def _retests_done(entry: dict) -> int: + """首探预算用 MAX_PROBES;超出部分算延迟回测次数。""" + return max(0, _probes(entry) - MAX_PROBES) + + +def should_probe(name: str, st: dict) -> bool: + """ + 新号:首探(≤MAX_PROBES 次当场连探)。 + 403 软状态:满首探后等 RETEST_AFTER_SEC,再给 RETEST_MAX 次回测。 + 429/真坏/终态:不探。 + """ + if name not in st: + return True + cur = st[name].get("status", "") + if cur == "discarded" or cur in TERMINAL_STATUSES: + return False + probes = _probes(st[name]) + + if cur in SOFT_RETEST_STATUSES: + if _retests_done(st[name]) >= RETEST_MAX: + return False + # 首探未打满(异常中断)→ 先补完首探,不算延迟回测 + if probes < MAX_PROBES: + return True + return _age_sec(st[name]) >= RETEST_AFTER_SEC + + # 新号 / pending / 临时错误:未满首探预算可探 + if probes < MAX_PROBES and cur in ("", "pending", "probe_err", "unknown"): + return True + return False + + +def process_file(name: str, st: dict, *, force_probe: bool = False) -> None: + src_path = AUTH_DIR / name + # 回测/整备也可只读 cpa_ready 已有 json(rescue 后 AUTH 可能仍在) + out_existing = OUT_DIR / name + if not src_path.exists() and not out_existing.exists(): + return + + src = None + if src_path.exists(): + src = load_source(src_path) + if (not src or not src.get("access_token")) and out_existing.exists(): + src = load_source(out_existing) + if not src: + return + if not src.get("access_token"): + log(f" ✗ {name} 无 access_token, 跳过") + return + + if name in st and st[name].get("status") == "discarded": + return + + sub, entry = finalize(src) + out_path = write_out(name, entry) + + already = _probes(st[name]) if name in st else 0 + prev_status = st[name].get("status", "") if name in st else "" + is_soft_retest = ( + prev_status in SOFT_RETEST_STATUSES + and already >= MAX_PROBES + and _retests_done(st.get(name, {})) < RETEST_MAX + ) + + if not force_probe and not should_probe(name, st): + if name in st and should_discard(st[name].get("status", ""), already): + discard_bad(name, st) + write_state(st) + return + # 软状态未到点:静默整备 + if prev_status in SOFT_RETEST_STATUSES: + left = max(0, RETEST_AFTER_SEC - _age_sec(st[name])) + if already >= MAX_PROBES and _retests_done(st[name]) < RETEST_MAX: + log( + f" • {name} 403 待回测 " + f"({left}s / retest={_retests_done(st[name])}/{RETEST_MAX})" + ) + return + log( + f" • {name} 整备覆盖 " + f"({prev_status or '-'}/probes={already}) → {out_path.name}" + ) + return + + if force_probe and already >= MAX_PROBES + RETEST_MAX and not os.environ.get("ACPA_FORCE_IGNORE_CAP"): + log(f" • {name} 已 probes={already},跳过强制重探(省额度)") + if should_discard(st.get(name, {}).get("status", ""), already): + discard_bad(name, st) + write_state(st) + return + + if is_soft_retest: + budget = 1 + phase = f"403回测×1 (after {RETEST_AFTER_SEC}s, done={_retests_done(st[name])}/{RETEST_MAX})" + else: + budget = max(0, MAX_PROBES - already) + phase = f"首探×≤{budget} (used={already}/{MAX_PROBES})" + + # 新号首探前 warmup:刚 mint 的 access 常被上游瞬时 403 + if not is_soft_retest and PROBE_WARMUP_SEC > 0: + time.sleep(PROBE_WARMUP_SEC) + + # 新号首探前 warmup:刚 mint 的 access 常被上游瞬时 403 + if not is_soft_retest and PROBE_WARMUP_SEC > 0: + time.sleep(PROBE_WARMUP_SEC) + + log(f" ▸ {name} 整备完成, {phase} ...") + r, used = probe_budgeted(out_path, budget) + status = r.get("status", "probe_err") + snip = r.get("snippet", "") + new_probes = already + used + + # 延迟回测仍 403 → 真坏 chat_dead + if is_soft_retest and status in ( + "chat_denied", "forbidden_domain_ban", "probe_cap", "probe_err", "unknown", + ): + status = "chat_dead" + log( + f" {name}: [chat_dead] 回测仍 403/失败 http={r.get('http')} " + f"probes={new_probes} {snip[:70]}" + ) + else: + log( + f" {name}: [{status}] http={r.get('http')} " + f"probes={new_probes} {snip[:70]}" + ) + + st[name] = { + "sub": sub, + "email": entry.get("email", ""), + "status": status, + "ts": int(time.time()), + "probes": new_probes, + } + if should_discard(status, new_probes): + discard_bad(name, st) + write_state(st) + + +def scan_once(st: dict) -> int: + """扫一轮:新号首探;403 到期回测;429 等真坏丢弃;alive 只刷整备。""" + if not AUTH_DIR.exists(): + log(f"⚠ 出货目录不存在: {AUTH_DIR} (auth-service.sh 启动后才会建)") + # 仍处理 cpa_ready 里待回测/rescue 的号 + pruned = prune_orphan_state(st) + if pruned: + log(f"本轮 prune 幽灵 state {pruned} 行") + discarded = sweep_discard(st) + if discarded: + log(f"本轮丢弃坏号 {discarded} 个 → {DISCARD_DIR.name}/") + + names: set[str] = set() + if AUTH_DIR.exists(): + names.update(p.name for p in AUTH_DIR.glob("xai-*.json")) + # cpa_ready 里所有现存 json 都纳入扫描: + # - soft 到期回测 + # - 无 state(删 TSV 后残留 / 手工丢入)→ 当新号首探 + # - alive 走 should_probe=False 分支只刷整备 + if OUT_DIR.exists(): + names.update(p.name for p in OUT_DIR.glob("xai-*.json")) + + n = 0 + for name in sorted(names): + if not RUN: + break + if name in st and st[name].get("status") == "discarded": + continue + if name in st and not should_probe(name, st): + if should_discard(st[name].get("status", ""), _probes(st[name])): + discard_bad(name, st) + write_state(st) + continue + src_path = AUTH_DIR / name if (AUTH_DIR / name).exists() else OUT_DIR / name + src = load_source(src_path) if src_path.exists() else None + if src and src.get("access_token"): + _, entry = finalize(src) + write_out(name, entry) + continue + process_file(name, st) + n += 1 + return n + + +def main() -> int: + ap = argparse.ArgumentParser(description="auto CPA 整备驻守") + ap.add_argument("--interval", type=float, default=2.0) + ap.add_argument("--once", action="store_true", help="扫一轮就退") + ap.add_argument( + "--reprobe-soft", + action="store_true", + help="立刻对 chat_denied/probe_cap 等软 403 做回测(仍 1 次/号,受 RETEST_MAX)", + ) + ap.add_argument( + "--no-rescue", + action="store_true", + help="启动时不从 _discarded 捞回误丢的 403", + ) + args = ap.parse_args() + + ensure_dir(OUT_DIR) + ensure_dir(DISCARD_DIR) + st = load_state() + n_prune = prune_orphan_state(st) + + # 历史 probe_cap 并入 soft,可走延迟回测 + for name, v in st.items(): + if v.get("status") == "probe_cap": + v["status"] = "chat_denied" + + soft_n = sum(1 for v in st.values() if v.get("status") in SOFT_RETEST_STATUSES) + bad_n = sum( + 1 for v in st.values() + if v.get("status") != "discarded" + and should_discard(v.get("status", ""), _probes(v)) + ) + log( + f"启动 | 出货={AUTH_DIR} | 落盘={OUT_DIR} | state={len(st)} 条 " + f"(启动 prune {n_prune}) | 首探≤{MAX_PROBES}/号 " + f"| warmup={PROBE_WARMUP_SEC}s 403即重试={IMMEDIATE_403_RETRIES}×{IMMEDIATE_403_SLEEP}s " + f"| 403回测 after={RETEST_AFTER_SEC}s ×{RETEST_MAX} " + f"| soft_403={soft_n} | pending_discard={bad_n}" + ) + + if not args.no_rescue: + n_res = rescue_soft_discards(st) + if n_res: + log(f"启动捞回误丢 403 {n_res} 个(将回测)") + + n0 = sweep_discard(st) + if n0: + log(f"启动丢弃真坏号 {n0} 个 → {DISCARD_DIR}/") + + if st and not args.reprobe_soft: + log( + f"(state 已有: 刷 alive 整备;403 满 {RETEST_AFTER_SEC}s 回测×{RETEST_MAX};" + f"429 free_exhausted 立刻丢)" + ) + for name in list(st.keys()): + if not RUN: + break + if st[name].get("status") != "alive": + continue + src_path = AUTH_DIR / name + if src_path.exists(): + src = load_source(src_path) + if src and src.get("access_token"): + _, entry = finalize(src) + write_out(name, entry) + + if args.reprobe_soft: + # soft 立刻回测一次:拨到期 + 把 probes 收成 MAX_PROBES, + # 避免旧版 probes=2 在新默认 MAX_PROBES=1 下被算成 retest 已用尽。 + # 只动 soft;alive 绝不复探(省 free 额度)。 + now = int(time.time()) + targets = [] + for name, v in st.items(): + if v.get("status") in SOFT_RETEST_STATUSES: + v["ts"] = now - RETEST_AFTER_SEC + v["probes"] = MAX_PROBES + targets.append(name) + write_state(st) + log(f"--reprobe-soft: {len(targets)} 个 403 软状态立刻回测×1(alive 不动)") + for name in targets: + if not RUN: + break + process_file(name, st, force_probe=True) + write_state(st) + alive = sum(1 for v in st.values() if v.get("status") == "alive") + dead = sum( + 1 for v in st.values() + if v.get("status") in DISCARD_STATUSES or v.get("status") == "discarded" + ) + soft = sum(1 for v in st.values() if v.get("status") in SOFT_RETEST_STATUSES) + log(f"回测结束 | alive={alive} soft_403={soft} dead/discard={dead} state={len(st)}") + return 0 + + while RUN: + try: + n = scan_once(st) + if n: + log(f"本轮处理 {n} 个号(新号/回测)") + except Exception as exc: + log(f"轮询异常: {type(exc).__name__} {exc}") + if args.once: + break + slept = 0.0 + while slept < args.interval and RUN: + time.sleep(0.25) + slept += 0.25 + + log(f"退出 | state={len(st)} 条 写 {STATE_FILE}") + write_state(st) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/keys/sync_acc.py b/keys/sync_acc.py new file mode 100755 index 0000000..fb8ef21 --- /dev/null +++ b/keys/sync_acc.py @@ -0,0 +1,188 @@ +#!/usr/bin/env python3 +""" +sync_acc.py — 把 keys/cpa_ready/ 里「探活通过」号的 access_token 同步进 keys/acc.md + +默认只入库 acpa_watchdog 记为 alive 的号(读 keys/cpa_ready/_state.tsv)。 +403 forbidden / token_dead / probe_err 等不进 acc.md,避免坏号灌库。 + +用法: + python3 sync_acc.py # 常驻轮询(默认 5s), Ctrl-C 退 + python3 sync_acc.py --once # 同步一次就退 + python3 sync_acc.py --all # 忽略探活状态,全量抽 token(排障用) + python3 sync_acc.py --interval 10 + +acc.md 每轮全量重写(去重),0600。 +""" +from __future__ import annotations + +import argparse +import glob +import json +import os +import signal +import sys +import time +from pathlib import Path + +PROJECT = Path(__file__).resolve().parent.parent # grok-free-register/ +SRC_DIR = PROJECT / "keys" / "cpa_ready" +STATE_FILE = SRC_DIR / "_state.tsv" +OUT_FILE = PROJECT / "keys" / "acc.md" + +# 默认可入库的探活状态(discarded/坏号不在 cpa_ready 主目录,也不在此集合) +ALIVE_STATUSES = frozenset({"alive"}) +SKIP_STATUSES = frozenset({ + "discarded", "free_exhausted", "rate_limited", "token_dead", + "no_quota_paid_end", "probe_cap", "chat_denied", "chat_dead", + "forbidden_domain_ban", "probe_err", "unknown", +}) + +RUN = True + + +def _sig(_s, _f): + global RUN + RUN = False + + +signal.signal(signal.SIGINT, _sig) +signal.signal(signal.SIGTERM, _sig) + + +def log(msg: str) -> None: + print(f"[sync_acc {time.strftime('%H:%M:%S')}] {msg}", flush=True) + + +def load_alive_names() -> set[str] | None: + """ + 读 _state.tsv → 状态为 alive 且盘上仍有 json 的文件名集合。 + 无 state 文件返回 None。幽灵 alive(TSV 有、json 无)不计。 + """ + if not STATE_FILE.exists(): + return None + alive: set[str] = set() + for ln in STATE_FILE.read_text(encoding="utf-8").splitlines(): + parts = ln.split("\t") + if len(parts) < 4: + continue + name, status = parts[0], parts[3] + if status in ALIVE_STATUSES: + # 只认盘上真实存在的号,避免删库后 state 幽灵把计数钉死 + if (SRC_DIR / name).exists(): + alive.add(name) + # discarded / 坏状态明确不进(双保险;主目录 json 也应已被 watchdog 挪走) + elif status in SKIP_STATUSES: + continue + return alive + + +def collect(*, require_alive: bool) -> tuple[list[str], dict[str, int]]: + """ + 返回 (去重 token 列表, 计数)。 + require_alive=True 时只收 _state.tsv 标为 alive 且文件仍在的 xai-*.json。 + """ + files = sorted(glob.glob(str(SRC_DIR / "xai-*.json"))) + alive_names = load_alive_names() if require_alive else None + + seen: set[str] = set() + out: list[str] = [] + stats = { + "files": len(files), + "alive_state": len(alive_names) if alive_names is not None else -1, + "skipped_not_alive": 0, + "no_tok": 0, + "bad": 0, + "no_state_skip_all": 0, + } + + if require_alive and alive_names is None: + # 有 cpa json 但还没 state:宁可不灌,避免 403 进库 + if files: + log(f"⚠ 无 {STATE_FILE.name},拒绝全量入库(等 watchdog 探活)") + stats["no_state_skip_all"] = len(files) + return [], stats + + for f in files: + name = Path(f).name + if require_alive and alive_names is not None and name not in alive_names: + stats["skipped_not_alive"] += 1 + continue + try: + d = json.loads(Path(f).read_text(encoding="utf-8")) + except Exception: + stats["bad"] += 1 + continue + tok = d.get("access_token") + if not isinstance(tok, str) or not tok: + stats["no_tok"] += 1 + continue + if tok in seen: + continue + seen.add(tok) + out.append(tok) + return out, stats + + +def write_acc(tokens: list[str]) -> None: + OUT_FILE.parent.mkdir(parents=True, exist_ok=True) + # 固定同目录临时文件,避免 with_suffix 边界情况 + 保证 replace 原子 + tmp = OUT_FILE.parent / (OUT_FILE.name + ".tmp") + tmp.write_text("\n".join(tokens) + ("\n" if tokens else ""), encoding="utf-8") + os.replace(tmp, OUT_FILE) + try: + os.chmod(OUT_FILE, 0o600) + except OSError: + pass + + +def sync_once(*, require_alive: bool) -> int: + toks, st = collect(require_alive=require_alive) + write_acc(toks) + mode = "alive-only" if require_alive else "all" + log( + f"sync → {OUT_FILE.name}: {len(toks)} tokens [{mode}] " + f"(files={st['files']} state_alive={st['alive_state']} " + f"skip_dead={st['skipped_not_alive']} no_tok={st['no_tok']} bad={st['bad']})" + ) + return len(toks) + + +def main() -> int: + ap = argparse.ArgumentParser( + description="cpa_ready access_token → keys/acc.md (default: alive only)" + ) + ap.add_argument("--interval", type=float, default=5.0) + ap.add_argument("--once", action="store_true", help="刷一次就退") + ap.add_argument( + "--all", + action="store_true", + help="忽略探活,全量抽 token(排障;会把 403 也写入 acc.md)", + ) + args = ap.parse_args() + require_alive = not args.all + + if not SRC_DIR.exists(): + log(f"⚠ 目录不存在: {SRC_DIR} (watchdog 整备后才会建)") + if args.once: + return 0 + + sync_once(require_alive=require_alive) + if args.once: + return 0 + + log(f"常驻轮询 interval={args.interval}s alive_only={require_alive} (Ctrl-C 退)") + while RUN: + slept = 0.0 + while slept < args.interval and RUN: + time.sleep(0.25) + slept += 0.25 + if not RUN: + break + sync_once(require_alive=require_alive) + + log("退出") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..237e2df --- /dev/null +++ b/requirements.txt @@ -0,0 +1,7 @@ +cloakbrowser>=0.3.0 +requests>=2.31.0 +python-dotenv>=1.0.0 +httpx>=0.28 +playwright>=1.55 +curl_cffi>=0.7.0 +curl_cffi>=0.7.0 diff --git a/reset_pipeline.sh b/reset_pipeline.sh new file mode 100755 index 0000000..69a0991 --- /dev/null +++ b/reset_pipeline.sh @@ -0,0 +1,222 @@ +#!/usr/bin/env bash +# 一键清零注册 → 认证 → CPA 整条流水线的运行时数据,方便重开。 +# +# 你只清 keys/acc.md / accounts.txt 不够: +# acpa_watchdog 会从 ~/Downloads/grok-free-register-auth/authenticated/ +# 把旧号再整备进 keys/cpa_ready/,sync_acc 再写回 acc.md。 +# +# 用法(本文件在项目根): +# bash reset_pipeline.sh # dry-run,只列将删内容 +# bash reset_pipeline.sh --yes # 真正删除 +# bash reset_pipeline.sh --yes --keep-register +# # 只清认证/CPA,保留 keys 里新注册的 SSO 源 +# bash reset_pipeline.sh --yes --keep-cpa +# # 只清认证账本/源快照,保留 cpa_ready(一般不需要) +# +# 建议先停 auth-service / register,再跑本脚本。 + +set -euo pipefail + +ROOT="$(cd "$(dirname "$0")" && pwd)" +KEYS="$ROOT/keys" +AUTH_DIR="${XAI_AUTH_SERVICE_LOCAL_DIR:-$HOME/Downloads/grok-free-register-auth}" + +YES=0 +KEEP_REGISTER=0 +KEEP_CPA=0 +KEEP_AUTH_SALT=1 + +for arg in "$@"; do + case "$arg" in + --yes|-y) YES=1 ;; + --keep-register) KEEP_REGISTER=1 ;; + --keep-cpa) KEEP_CPA=1 ;; + --wipe-salt) KEEP_AUTH_SALT=0 ;; + -h|--help) + sed -n '2,20p' "$0" + exit 0 + ;; + *) + echo "未知参数: $arg" >&2 + exit 2 + ;; + esac +done + +# 运行时产物清单(脚本本身永不删) +REGISTER_FILES=( + "$KEYS/accounts.txt" + "$KEYS/grok.txt" + "$KEYS/auth-sessions.jsonl" + "$KEYS/acc.md" +) +CPA_GLOBS=( + "$KEYS/cpa_ready/xai-"*.json + "$KEYS/cpa_ready/_state.tsv" +) +AUTH_PATHS=( + "$AUTH_DIR/authenticated" + "$AUTH_DIR/claimed" + "$AUTH_DIR/enrollment-ledger.db" + "$AUTH_DIR/source-snapshot.jsonl" +) +AUTH_SALT="$AUTH_DIR/.ledger-salt" + +count_matches() { + local n=0 + local p + for p in "$@"; do + # shellcheck disable=SC2086 + if compgen -G "$p" > /dev/null 2>&1; then + while IFS= read -r -d '' f; do + n=$((n + 1)) + done < <(find $p -type f -print0 2>/dev/null || true) + # also count plain files matched by glob/file + if [[ -f "$p" ]]; then + n=$((n + 1)) + fi + elif [[ -f "$p" ]]; then + n=$((n + 1)) + elif [[ -d "$p" ]]; then + while IFS= read -r -d '' f; do + n=$((n + 1)) + done < <(find "$p" -type f -print0 2>/dev/null || true) + fi + done + echo "$n" +} + +size_of() { + local p="$1" + if [[ -e "$p" ]]; then + du -sh "$p" 2>/dev/null | awk '{print $1}' + else + echo "-" + fi +} + +echo "=== grok-free-register 流水线清零 ===" +echo "ROOT = $ROOT" +echo "KEYS = $KEYS" +echo "AUTH_DIR = $AUTH_DIR" +echo + +# 检测可能还在跑的服务 +if pgrep -af 'xai_enroller|acpa_watchdog|sync_acc|grok_register.register' 2>/dev/null | grep -v "reset_pipeline\|pgrep\|grep" >/dev/null; then + echo "[!] 检测到相关进程仍在运行:" + pgrep -af 'xai_enroller|acpa_watchdog|sync_acc|grok_register.register' 2>/dev/null | grep -v "reset_pipeline\|pgrep" || true + echo " 建议先 Ctrl-C 停掉 auth-service / register,再 --yes" + echo +fi + +echo "-- 将处理的目标 --" +if [[ "$KEEP_REGISTER" -eq 0 ]]; then + for f in "${REGISTER_FILES[@]}"; do + if [[ -e "$f" ]]; then + echo " [register] $(size_of "$f") $f" + fi + done +else + echo " [register] 跳过 (--keep-register)" +fi + +if [[ "$KEEP_CPA" -eq 0 ]]; then + cpa_n=$(find "$KEYS/cpa_ready" -name 'xai-*.json' -type f 2>/dev/null | wc -l | tr -d ' ') + echo " [cpa] ${cpa_n} 个 xai-*.json @ $KEYS/cpa_ready/" + [[ -f "$KEYS/cpa_ready/_state.tsv" ]] && echo " [cpa] _state.tsv" +else + echo " [cpa] 跳过 (--keep-cpa)" +fi + +auth_n=$(find "$AUTH_DIR/authenticated" -type f 2>/dev/null | wc -l | tr -d ' ') +echo " [auth] authenticated/ ${auth_n} 文件 ($(size_of "$AUTH_DIR/authenticated"))" +[[ -d "$AUTH_DIR/claimed" ]] && echo " [auth] claimed/ $(find "$AUTH_DIR/claimed" -type f 2>/dev/null | wc -l | tr -d ' ') 文件" +[[ -f "$AUTH_DIR/enrollment-ledger.db" ]] && echo " [auth] enrollment-ledger.db $(size_of "$AUTH_DIR/enrollment-ledger.db")" +[[ -f "$AUTH_DIR/source-snapshot.jsonl" ]] && echo " [auth] source-snapshot.jsonl $(size_of "$AUTH_DIR/source-snapshot.jsonl") lines=$(wc -l < "$AUTH_DIR/source-snapshot.jsonl" 2>/dev/null | tr -d ' ')" +if [[ "$KEEP_AUTH_SALT" -eq 0 && -f "$AUTH_SALT" ]]; then + echo " [auth] .ledger-salt (将删,指纹会重算)" +else + echo " [auth] .ledger-salt 保留" +fi +echo + +if [[ "$YES" -ne 1 ]]; then + echo "dry-run 模式,未删除任何东西。" + echo "确认后执行: bash reset_pipeline.sh --yes" + exit 0 +fi + +rm_file() { + local f="$1" + if [[ -f "$f" || -L "$f" ]]; then + rm -f -- "$f" + echo " removed file $f" + fi +} + +wipe_dir_contents() { + local d="$1" + if [[ -d "$d" ]]; then + # 只清内容,保留目录本身 + find "$d" -mindepth 1 -maxdepth 1 -exec rm -rf -- {} + + echo " wiped dir $d/" + fi +} + +echo "-- 执行删除 --" + +if [[ "$KEEP_REGISTER" -eq 0 ]]; then + for f in "${REGISTER_FILES[@]}"; do + rm_file "$f" + done + # 重建空占位,避免下游 open 报错 + : > "$KEYS/accounts.txt" + : > "$KEYS/grok.txt" + : > "$KEYS/auth-sessions.jsonl" + : > "$KEYS/acc.md" + chmod 600 "$KEYS/acc.md" "$KEYS/auth-sessions.jsonl" 2>/dev/null || true + echo " recreated empty accounts.txt / grok.txt / auth-sessions.jsonl / acc.md" +fi + +if [[ "$KEEP_CPA" -eq 0 ]]; then + mkdir -p "$KEYS/cpa_ready" + find "$KEYS/cpa_ready" -type f \( -name 'xai-*.json' -o -name '_state.tsv' \) -delete 2>/dev/null || true + echo " wiped $KEYS/cpa_ready/xai-*.json + _state.tsv" +fi + +# 认证侧:这是旧号回灌的真正源头 +mkdir -p "$AUTH_DIR/authenticated" "$AUTH_DIR/claimed" +wipe_dir_contents "$AUTH_DIR/authenticated" +wipe_dir_contents "$AUTH_DIR/claimed" +rm_file "$AUTH_DIR/enrollment-ledger.db" +rm_file "$AUTH_DIR/source-snapshot.jsonl" +# 偶发的 sqlite 旁路文件 +rm_file "$AUTH_DIR/enrollment-ledger.db-wal" +rm_file "$AUTH_DIR/enrollment-ledger.db-shm" +rm_file "$AUTH_DIR/enrollment-ledger.db-journal" + +if [[ "$KEEP_AUTH_SALT" -eq 0 ]]; then + rm_file "$AUTH_SALT" +fi + +# 兜底:项目内若误放了同名目录也清 +for extra in "$KEYS/authenticated" "$KEYS/claimed"; do + if [[ -d "$extra" ]]; then + wipe_dir_contents "$extra" + fi +done + +echo +echo "-- 清零后状态 --" +echo " accounts.txt lines = $(wc -l < "$KEYS/accounts.txt" 2>/dev/null | tr -d ' ')" +echo " auth-sessions = $(wc -l < "$KEYS/auth-sessions.jsonl" 2>/dev/null | tr -d ' ')" +echo " acc.md tokens = $(wc -l < "$KEYS/acc.md" 2>/dev/null | tr -d ' ')" +echo " cpa_ready json = $(find "$KEYS/cpa_ready" -name 'xai-*.json' -type f 2>/dev/null | wc -l | tr -d ' ')" +echo " authenticated json = $(find "$AUTH_DIR/authenticated" -type f 2>/dev/null | wc -l | tr -d ' ')" +echo " ledger.db = $([[ -f "$AUTH_DIR/enrollment-ledger.db" ]] && echo present || echo gone)" +echo " snapshot = $([[ -f "$AUTH_DIR/source-snapshot.jsonl" ]] && echo present || echo gone)" +echo +echo "完成。重开顺序建议:" +echo " 1) bash start.sh # 注册出 SSO → keys/auth-sessions.jsonl + accounts.txt" +echo " 2) bash auth-service.sh # device-flow → AUTH_DIR/authenticated + acpa→cpa_ready + sync_acc" +echo " 若只想清 CPA/认证、保留刚注的 SSO: 下次加 --keep-register" diff --git a/scripts/clean_history.sh b/scripts/clean_history.sh new file mode 100755 index 0000000..df94be7 --- /dev/null +++ b/scripts/clean_history.sh @@ -0,0 +1,140 @@ +#!/usr/bin/env bash +# 快速清历史归档(不碰现网号池)。 +# +# 只删这些「历史」: +# keys/cpa_ready/cpa_ready_*.zip 导出包 +# keys/cpa_ready/_discarded/** 已丢弃号 +# keys/*.bak* / keys/*~ / keys/__pycache__ +# keys/async_auth.log / logs/* +# keys/212.zip 等杂包 +# +# 默认保留: +# keys/cpa_ready/xai-*.json +# keys/cpa_ready/_state.tsv +# keys/acc.md / accounts.txt / auth-sessions.jsonl +# AUTH_DIR/authenticated 现网出货 +# +# 用法(项目根): +# bash scripts/clean_history.sh # dry-run +# bash scripts/clean_history.sh --yes # 真删 +# bash scripts/clean_history.sh --yes --deep # 再清 source-snapshot / 旧 .bak 脚本 +# +# 全量清零号池请用: bash reset_pipeline.sh --yes + +set -euo pipefail + +ROOT="$(cd "$(dirname "$0")/.." && pwd)" +KEYS="$ROOT/keys" +CPA="$KEYS/cpa_ready" +AUTH_DIR="${XAI_AUTH_SERVICE_LOCAL_DIR:-${XAI_ENROLLER_LOCAL_AUTH_DIR:-$HOME/Downloads/grok-free-register-auth}}" +LOGS="$ROOT/logs" + +YES=0 +DEEP=0 + +for arg in "$@"; do + case "$arg" in + --yes|-y) YES=1 ;; + --deep) DEEP=1 ;; + -h|--help) + sed -n '2,24p' "$0" + exit 0 + ;; + *) + echo "未知参数: $arg" >&2 + exit 2 + ;; + esac +done + +size_of() { + local p="$1" + if [[ -e "$p" ]]; then + du -sh "$p" 2>/dev/null | awk '{print $1}' + else + echo "-" + fi +} + +list_targets() { + # zips / archives under cpa_ready and keys + find "$CPA" -maxdepth 1 -type f \( -name 'cpa_ready_*.zip' -o -name '*.zip' -o -name '*.tar' -o -name '*.tar.gz' -o -name '*.tgz' \) 2>/dev/null || true + find "$KEYS" -maxdepth 1 -type f \( -name '*.zip' -o -name '*.tar' -o -name '*.tar.gz' -o -name '*.tgz' \) 2>/dev/null || true + # discarded + if [[ -d "$CPA/_discarded" ]]; then + find "$CPA/_discarded" -mindepth 1 2>/dev/null || true + fi + # logs / pyc / bak + find "$KEYS" -maxdepth 1 -type f \( -name '*.log' -o -name '*.bak' -o -name '*.bak.*' -o -name '*~' \) 2>/dev/null || true + find "$KEYS" -maxdepth 1 -type d -name '__pycache__' 2>/dev/null || true + if [[ -d "$LOGS" ]]; then + find "$LOGS" -type f 2>/dev/null || true + fi + if [[ "$DEEP" -eq 1 ]]; then + # deep: old snapshot growth + script backups; still keep live pool + [[ -f "$AUTH_DIR/source-snapshot.jsonl" ]] && echo "$AUTH_DIR/source-snapshot.jsonl" + find "$KEYS" -maxdepth 1 -type f -name 'acpa_watchdog.py.bak*' 2>/dev/null || true + fi +} + +mapfile -t TARGETS < <(list_targets | awk 'NF' | sort -u) + +echo "=== grok-free-register 快速清历史 ===" +echo "ROOT = $ROOT" +echo "KEYS = $KEYS" +echo "AUTH_DIR = $AUTH_DIR" +echo "MODE = $([[ $YES -eq 1 ]] && echo APPLY || echo dry-run)$([[ $DEEP -eq 1 ]] && echo ' +deep' || true)" +echo + +if [[ "${#TARGETS[@]}" -eq 0 ]]; then + echo "没有可清的历史文件。" + exit 0 +fi + +echo "-- 将处理 --" +total=0 +for f in "${TARGETS[@]}"; do + if [[ -e "$f" ]]; then + echo " $(size_of "$f") $f" + total=$((total + 1)) + fi +done +echo " 共 $total 项" +echo + +# live pool summary (never touched) +echo "-- 现网号池(保留)--" +echo " cpa xai-*.json = $(find "$CPA" -maxdepth 1 -name 'xai-*.json' -type f 2>/dev/null | wc -l | tr -d ' ')" +echo " acc.md lines = $(wc -l < "$KEYS/acc.md" 2>/dev/null | tr -d ' ' || echo 0)" +echo " sessions = $(wc -l < "$KEYS/auth-sessions.jsonl" 2>/dev/null | tr -d ' ' || echo 0)" +echo " authenticated = $(find "$AUTH_DIR/authenticated" -type f 2>/dev/null | wc -l | tr -d ' ' || echo 0)" +echo + +if [[ "$YES" -ne 1 ]]; then + echo "dry-run。确认后: bash scripts/clean_history.sh --yes" + echo "更深一层(含 snapshot): bash scripts/clean_history.sh --yes --deep" + exit 0 +fi + +echo "-- 执行删除 --" +for f in "${TARGETS[@]}"; do + if [[ -d "$f" ]]; then + rm -rf -- "$f" + echo " removed dir $f" + elif [[ -e "$f" ]]; then + rm -f -- "$f" + echo " removed file $f" + fi +done + +# keep empty discarded dir +mkdir -p "$CPA/_discarded" 2>/dev/null || true +mkdir -p "$LOGS" 2>/dev/null || true + +echo +echo "-- 清后 --" +echo " cpa zip left = $(find "$CPA" -maxdepth 1 -name '*.zip' -type f 2>/dev/null | wc -l | tr -d ' ')" +echo " discarded files= $(find "$CPA/_discarded" -type f 2>/dev/null | wc -l | tr -d ' ')" +echo " cpa xai-*.json = $(find "$CPA" -maxdepth 1 -name 'xai-*.json' -type f 2>/dev/null | wc -l | tr -d ' ')" +echo " keys size = $(size_of "$KEYS")" +echo "完成。" diff --git a/scripts/ensure_runtime.sh b/scripts/ensure_runtime.sh new file mode 100755 index 0000000..072dd25 --- /dev/null +++ b/scripts/ensure_runtime.sh @@ -0,0 +1,32 @@ +#!/usr/bin/env bash + +ensure_runtime() { + local lock_dir=".setup.lock" + local acquired=0 + local attempt + + for attempt in {1..300}; do + if mkdir "$lock_dir" 2>/dev/null; then + acquired=1 + break + fi + sleep 0.2 + done + if [ "$acquired" -ne 1 ]; then + echo "[!] 另一个安装进程长时间未结束,请稍后重试。" >&2 + return 1 + fi + trap 'rmdir .setup.lock 2>/dev/null || true' EXIT INT TERM + + if [ ! -d .venv ]; then + echo "[*] 首次运行,安装依赖..." + if ! bash setup.sh; then + rmdir "$lock_dir" + trap - EXIT INT TERM + return 1 + fi + fi + + rmdir "$lock_dir" + trap - EXIT INT TERM +} diff --git a/scripts/export_registered_sessions.py b/scripts/export_registered_sessions.py new file mode 100755 index 0000000..beda0ae --- /dev/null +++ b/scripts/export_registered_sessions.py @@ -0,0 +1,206 @@ +#!/usr/bin/env python3 +"""通过 SSH 导出不含密码的注册会话;兼容历史裸 SSO 记录。""" + +import argparse +import json +import os +import sys +import time +from pathlib import Path + + +COOKIE_FIELDS = frozenset( + { + "name", + "value", + "url", + "domain", + "path", + "expires", + "httpOnly", + "secure", + "sameSite", + } +) +MAX_RECORD_BYTES = 256 * 1024 +LEGACY_COOKIE_SCOPE = { + "name": "sso", + "domain": "accounts.x.ai", + "path": "/", + "secure": True, + "httpOnly": True, + "sameSite": "Lax", +} + + +def _decode_json_line(raw, label): + if len(raw) > MAX_RECORD_BYTES: + raise ValueError(f"invalid {label} record") + try: + document = json.loads(raw.decode("utf-8")) + email = document["email"] + cookies = document["cookies"] + except (UnicodeDecodeError, TypeError, ValueError, KeyError) as exc: + raise ValueError(f"invalid {label} record") from exc + if not isinstance(email, str) or not email or not isinstance(cookies, list) or not cookies: + raise ValueError(f"invalid {label} record") + try: + email.encode("utf-8") + except UnicodeEncodeError as exc: + raise ValueError(f"invalid {label} record") from exc + normalized = [] + for cookie in cookies: + if not isinstance(cookie, dict): + raise ValueError(f"invalid {label} record") + filtered = {key: cookie[key] for key in COOKIE_FIELDS if key in cookie} + if not all( + isinstance(filtered.get(key), str) and filtered[key] + for key in ("name", "value") + ): + raise ValueError(f"invalid {label} record") + scope = filtered.get("domain") or filtered.get("url") + if not isinstance(scope, str) or not scope: + raise ValueError(f"invalid {label} record") + try: + filtered["name"].encode("utf-8") + filtered["value"].encode("utf-8") + scope.encode("utf-8") + except UnicodeEncodeError as exc: + raise ValueError(f"invalid {label} record") from exc + normalized.append(filtered) + return {"email": email, "cookies": normalized} + + +def _complete_lines(data): + """Return newline-terminated records and the unconsumed trailing bytes.""" + parts = data.split(b"\n") + return parts[:-1], parts[-1] + + +def _read_complete_file(path): + if not path.exists(): + return [] + lines, _incomplete = _complete_lines(path.read_bytes()) + return [line for line in lines if line] + + +def load_snapshots(path, *, raw_lines=None): + snapshots = {} + scopes = {} + lines = _read_complete_file(path) if raw_lines is None else raw_lines + for line in lines: + document = _decode_json_line(line, "session") + email = document["email"] + cookies = document["cookies"] + if email in snapshots: + continue + snapshots[email] = {"email": email, "cookies": cookies} + for cookie in cookies: + name = cookie.get("name") + domain = cookie.get("domain") + if name in {"sso", "sso-rw"} and domain: + scopes[(name, domain)] = { + "name": name, + "domain": domain, + "path": cookie.get("path", "/"), + "secure": bool(cookie.get("secure", True)), + "httpOnly": bool(cookie.get("httpOnly", True)), + "sameSite": cookie.get("sameSite", "Lax"), + } + return snapshots, scopes + + +def export_sessions(sessions_path, accounts_path, *, session_lines=None): + snapshots, scopes = load_snapshots(sessions_path, raw_lines=session_lines) + for document in snapshots.values(): + yield document + if not accounts_path.exists(): + return + legacy_scopes = list(scopes.values()) or [LEGACY_COOKIE_SCOPE] + for raw in _read_complete_file(accounts_path): + try: + email, _password, sso = raw.decode("utf-8").rsplit(":", 2) + except (UnicodeDecodeError, ValueError) as exc: + raise ValueError("invalid account record") from exc + if not email or not sso or email in snapshots: + continue + cookies = [{**scope, "value": sso} for scope in legacy_scopes] + yield {"email": email, "cookies": cookies} + + +def _write_document(document): + payload = json.dumps(document, separators=(",", ":")) + if len(payload.encode("utf-8")) > MAX_RECORD_BYTES: + raise ValueError("invalid session record") + print(payload, flush=False) + + +def export_and_follow(sessions_path, accounts_path, *, poll_seconds=0.25): + """Emit a complete snapshot, then losslessly follow the same JSONL fd.""" + sessions_path.parent.mkdir(parents=True, exist_ok=True) + if not sessions_path.exists(): + fd = os.open(sessions_path, os.O_CREAT | os.O_EXCL | os.O_WRONLY, 0o600) + os.close(fd) + os.chmod(sessions_path, 0o600) + stream = sessions_path.open("rb") + + with stream: + opened = os.fstat(stream.fileno()) + initial_lines, pending = _complete_lines(stream.read()) + if len(pending) > MAX_RECORD_BYTES: + raise ValueError("invalid session record") + for document in export_sessions( + sessions_path, + accounts_path, + session_lines=[line for line in initial_lines if line], + ): + _write_document(document) + sys.stdout.flush() + + while True: + chunk = stream.read() + if chunk: + complete, pending = _complete_lines(pending + chunk) + if len(pending) > MAX_RECORD_BYTES: + raise ValueError("invalid session record") + for raw in complete: + if raw: + _write_document(_decode_json_line(raw, "session")) + sys.stdout.flush() + continue + + try: + current = sessions_path.stat() + except FileNotFoundError: + return 3 + if ( + current.st_dev != opened.st_dev + or current.st_ino != opened.st_ino + or current.st_size < stream.tell() + ): + return 3 + time.sleep(poll_seconds) + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--follow", action="store_true") + parser.add_argument("sessions_path", type=Path) + parser.add_argument("accounts_path", type=Path) + args = parser.parse_args() + try: + if args.follow: + raise SystemExit(export_and_follow(args.sessions_path, args.accounts_path)) + for document in export_sessions(args.sessions_path, args.accounts_path): + _write_document(document) + except ValueError: + raise SystemExit(4) from None + except BrokenPipeError: + try: + sys.stdout.close() + finally: + raise SystemExit(0) + + +if __name__ == "__main__": + main() diff --git a/scripts/push_keys_to_auth.sh b/scripts/push_keys_to_auth.sh new file mode 100755 index 0000000..3c7adbf --- /dev/null +++ b/scripts/push_keys_to_auth.sh @@ -0,0 +1,153 @@ +#!/usr/bin/env bash +# 本机注册 → 推 SSO 源到远端 auth 机(local 模式吃 keys/) +# +# 动态公网 IP 不适合「远端 SSH 拉本机」;改由本机主动推: +# keys/auth-sessions.jsonl +# keys/accounts.txt +# → AUTH_HOST:REMOTE_ROOT/keys/ +# +# 用法(本机项目根或任意 cwd): +# bash scripts/push_keys_to_auth.sh # 推一次 +# bash scripts/push_keys_to_auth.sh --watch # 常驻,文件变了再推 +# bash scripts/push_keys_to_auth.sh --interval 15 +# +# 环境变量(可写本机 .env,本脚本会 source): +# AUTH_SSH_HOST=user@auth-server.example +# AUTH_REMOTE_ROOT=/opt/grok-free-register +# AUTH_SSH_IDENTITY= # 可选 -i 路径 +# AUTH_PUSH_INTERVAL=20 # --watch 轮询秒 +# +# 无内置默认主机/路径:必须显式设置 AUTH_SSH_HOST 与 AUTH_REMOTE_ROOT。 + +set -euo pipefail +ROOT="$(cd "$(dirname "$0")/.." && pwd)" +cd "$ROOT" + +# 调用方已 export 的 AUTH_* 优先(.env 不得盖掉),方便并行推多台 auth +_PRE_AUTH_SSH_HOST="${AUTH_SSH_HOST-}" +_PRE_AUTH_REMOTE_ROOT="${AUTH_REMOTE_ROOT-}" +_PRE_AUTH_SSH_IDENTITY="${AUTH_SSH_IDENTITY-}" +_PRE_AUTH_PUSH_INTERVAL="${AUTH_PUSH_INTERVAL-}" + +if [[ -f "$ROOT/.env" ]]; then + set -a + # shellcheck disable=SC1091 + . "$ROOT/.env" + set +a +fi + +AUTH_SSH_HOST="${_PRE_AUTH_SSH_HOST:-${AUTH_SSH_HOST:-}}" +AUTH_REMOTE_ROOT="${_PRE_AUTH_REMOTE_ROOT:-${AUTH_REMOTE_ROOT:-}}" +AUTH_SSH_IDENTITY="${_PRE_AUTH_SSH_IDENTITY:-${AUTH_SSH_IDENTITY:-}}" +AUTH_PUSH_INTERVAL="${_PRE_AUTH_PUSH_INTERVAL:-${AUTH_PUSH_INTERVAL:-20}}" + +if [[ -z "$AUTH_SSH_HOST" || -z "$AUTH_REMOTE_ROOT" ]]; then + echo "push_keys_to_auth: set AUTH_SSH_HOST and AUTH_REMOTE_ROOT (env or .env)" >&2 + echo " e.g. AUTH_SSH_HOST=user@auth-server.example AUTH_REMOTE_ROOT=/opt/grok-free-register" >&2 + exit 2 +fi + +WATCH=0 +INTERVAL="$AUTH_PUSH_INTERVAL" +for a in "$@"; do + case "$a" in + --watch) WATCH=1 ;; + --interval) + # next arg handled below if present as --interval=N form preferred + ;; + --interval=*) INTERVAL="${a#--interval=}" ;; + -h|--help) + sed -n '2,22p' "$0" + exit 0 + ;; + esac +done +# support: --interval 15 +prev="" +for a in "$@"; do + if [[ "$prev" == "--interval" ]]; then + INTERVAL="$a" + fi + prev="$a" +done + +SSH_OPTS=(-o BatchMode=yes -o ConnectTimeout=15 -o ServerAliveInterval=15) +if [[ -n "$AUTH_SSH_IDENTITY" ]]; then + SSH_OPTS+=(-i "$AUTH_SSH_IDENTITY") +fi + +KEYS_SRC="$ROOT/keys" +REMOTE_KEYS="${AUTH_REMOTE_ROOT}/keys" +# 只推 auth 源,不推 cpa_ready/acc(远端自己整备) +FILES=(auth-sessions.jsonl accounts.txt) + +log() { printf '[push_keys %s] %s\n' "$(date +%H:%M:%S)" "$*"; } + +fingerprint() { + local f + for f in "${FILES[@]}"; do + if [[ -f "$KEYS_SRC/$f" ]]; then + # size + mtime + sha256 head — cheap change detect + stat -c '%s %Y' "$KEYS_SRC/$f" 2>/dev/null || stat -f '%z %m' "$KEYS_SRC/$f" + sha256sum "$KEYS_SRC/$f" 2>/dev/null | awk '{print $1}' + else + echo "missing:$f" + fi + done +} + +push_once() { + local missing=0 f + for f in "${FILES[@]}"; do + if [[ ! -f "$KEYS_SRC/$f" ]]; then + log "⚠ 缺 $KEYS_SRC/$f" + missing=1 + fi + done + if [[ "$missing" -eq 1 ]]; then + return 1 + fi + + # 远端原子替换:先推 .push.tmp 再 mv + ssh "${SSH_OPTS[@]}" "$AUTH_SSH_HOST" "mkdir -p $(printf %q "$REMOTE_KEYS") && chmod 700 $(printf %q "$REMOTE_KEYS") 2>/dev/null || true" + + local remote_tmp remote_final + for f in "${FILES[@]}"; do + remote_tmp="${REMOTE_KEYS}/.${f}.push.tmp" + remote_final="${REMOTE_KEYS}/${f}" + # rsync over ssh when available; fallback scp + if command -v rsync >/dev/null 2>&1; then + rsync -az -e "ssh ${SSH_OPTS[*]}" \ + "$KEYS_SRC/$f" \ + "${AUTH_SSH_HOST}:${remote_tmp}" + else + scp "${SSH_OPTS[@]}" "$KEYS_SRC/$f" "${AUTH_SSH_HOST}:${remote_tmp}" + fi + ssh "${SSH_OPTS[@]}" "$AUTH_SSH_HOST" \ + "chmod 600 $(printf %q "$remote_tmp") && mv -f $(printf %q "$remote_tmp") $(printf %q "$remote_final")" + done + + local n_sess n_acc + n_sess=$(wc -l < "$KEYS_SRC/auth-sessions.jsonl" | tr -d ' ') + n_acc=$(wc -l < "$KEYS_SRC/accounts.txt" | tr -d ' ') + log "→ ${AUTH_SSH_HOST}:${REMOTE_KEYS}/ sessions=${n_sess} accounts=${n_acc}" +} + +LAST_FP="" +if [[ "$WATCH" -eq 0 ]]; then + push_once + exit $? +fi + +log "watch host=${AUTH_SSH_HOST} root=${AUTH_REMOTE_ROOT} interval=${INTERVAL}s" +while true; do + FP="$(fingerprint)" + if [[ "$FP" != "$LAST_FP" ]]; then + if push_once; then + LAST_FP="$FP" + else + log "推送失败,${INTERVAL}s 后重试" + fi + fi + sleep "$INTERVAL" +done diff --git a/setup.sh b/setup.sh new file mode 100755 index 0000000..0d63f04 --- /dev/null +++ b/setup.sh @@ -0,0 +1,50 @@ +#!/bin/bash +# Grok Free Register — 一键安装脚本 +# 用法: bash setup.sh + +set -e + +echo "=== Grok Free Register 安装 ===" + +# 检测系统 +if [ -f /etc/debian_version ]; then + echo "[1/4] 安装系统依赖 (Debian/Ubuntu)..." + sudo apt update -qq + sudo apt install -y -qq \ + python3 python3-pip python3-venv \ + libatk1.0-0t64 libatk-bridge2.0-0t64 libcups2t64 \ + libdrm2 libxkbcommon0 libxcomposite1 libxdamage1 \ + libxfixes3 libxrandr2 libgbm1 libpango-1.0-0 \ + libcairo2 libasound2t64 libnspr4 libnss3 libxshmfence1 \ + 2>/dev/null || true + # 兼容旧版 Ubuntu + sudo apt install -y -qq libatk1.0-0 libatk-bridge2.0-0 libcups2 libasound2 2>/dev/null || true +elif [ -f /etc/redhat-release ]; then + echo "[1/4] 安装系统依赖 (RHEL/CentOS)..." + sudo yum install -y -q \ + python3 python3-pip \ + atk cups-libs libdrm libXcomposite libXdamage libXfixes libXrandr \ + mesa-libgbm pango cairo alsa-lib nspr nss libxshmfence \ + 2>/dev/null || true +else + echo "[1/4] 未知系统,跳过系统依赖(如 Chrome 启动失败请手动安装)" +fi + +# Python 虚拟环境 +echo "[2/4] 创建 Python 环境..." +python3 -m venv .venv +.venv/bin/pip install -q --upgrade pip +.venv/bin/pip install -q -r requirements.txt + +# 本项目直接由 Playwright 启动二进制,因此显式准备 CloakBrowser Chromium。 +echo "[3/4] 下载 CloakBrowser Chromium..." +.venv/bin/python -m cloakbrowser install + +# 创建输出目录 +mkdir -p keys + +echo "[4/4] 安装完成!" +echo "" +echo "运行注册服务: bash start.sh" +echo "运行自建邮箱服务: bash start.sh --email-service" +echo "运行认证服务: bash auth-service.sh" diff --git a/start.sh b/start.sh new file mode 100755 index 0000000..76e4e80 --- /dev/null +++ b/start.sh @@ -0,0 +1,80 @@ +#!/bin/bash +# 一键启动:自动装依赖 → 引导配置 → 运行 +# 用法: +# bash start.sh # 首次会引导选模式,之后直接启动 +# bash start.sh --reconfig # 重新选择邮箱模式 +# bash start.sh --debug # 保留完整调试面板 +set -e +cd "$(dirname "$0")" + +. scripts/ensure_runtime.sh +ensure_runtime + +if [ "${1:-}" = "--email-service" ]; then + shift + if command -v flock >/dev/null 2>&1; then + mkdir -p logs + exec 8>logs/email-service.lock + if ! flock -n 8; then + echo "[!] 邮箱服务已经在运行。" + exit 1 + fi + fi + echo "[*] 启动邮箱服务... (Ctrl-C 停止)" + exec .venv/bin/python -m grok_register.email_server "$@" +fi + +reconfig=0 +register_args=() +for arg in "$@"; do + if [ "$arg" = "--reconfig" ]; then + reconfig=1 + else + register_args+=("$arg") + fi +done + +# 同一工作目录只允许一个注册进程,避免重复启动同时写账号和日志。 +if command -v flock >/dev/null 2>&1; then + mkdir -p logs + exec 9>logs/register.lock + if ! flock -n 9; then + echo "[!] 注册机已经在运行。" + exit 1 + fi +fi + +# 1) 配置:无 .env 或显式 --reconfig 时进入引导 +if [ ! -f .env ] || [ "$reconfig" -eq 1 ]; then + echo "" + echo "选择邮箱模式:" + echo " [1] 免费临时邮箱 (默认 · 零配置 · 直接回车 · 多 provider 自动 fallback)" + echo " [2] 自建域名邮箱 (需 Cloudflare Email Routing + 本地 webhook)" + read -rp "输入 1 或 2 [1]: " mode || mode=1 + if [ "$mode" = "2" ]; then + read -rp " 你的域名 (如 example.com): " domain + read -rp " webhook 地址 [http://127.0.0.1:8080]: " api + api=${api:-http://127.0.0.1:8080} + cat > .env < .env + fi + echo "[*] 已写入 .env" +fi + +# 2) 运行 +echo "[*] 启动注册服务... (Ctrl-C 停止)" +exec .venv/bin/python -m grok_register.register "${register_args[@]}" diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..dc6c822 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,92 @@ +""" +共享 fixtures — 供四层测试共用 +""" +import sys +import os +import asyncio +import pytest + +# 确保项目根目录在 sys.path 中 +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from grok_register.core.inventory import Inventory +from grok_register.core.observer import Metrics +from tests.fakes import ( + FakeTurnstile, FakeEmailService, FakeRegisterAPI, + Conservation, EventLog, +) + + +@pytest.fixture +def metrics(): + return Metrics() + + +@pytest.fixture +def inventory(metrics): + return Inventory(metrics=metrics) + + +@pytest.fixture +def sems(): + """标准 Semaphore 集合(小容量,便于测试)。""" + return { + 'physical': asyncio.Semaphore(4), + 't_slot': asyncio.Semaphore(4), + 'q_slot': asyncio.Semaphore(4), + 'q_pending': asyncio.Semaphore(6), + } + + +@pytest.fixture +def large_sems(): + """大容量 Semaphore(压力测试用)。""" + return { + 'physical': asyncio.Semaphore(32), + 't_slot': asyncio.Semaphore(64), + 'q_slot': asyncio.Semaphore(64), + 'q_pending': asyncio.Semaphore(48), + } + + +@pytest.fixture +def conservation(sems, inventory): + return Conservation( + sems['t_slot'], sems['q_slot'], sems['q_pending'], inventory + ) + + +@pytest.fixture +def fake_turnstile(): + return FakeTurnstile() + + +@pytest.fixture +def fake_email(): + return FakeEmailService() + + +@pytest.fixture +def fake_register(): + return FakeRegisterAPI() + + +@pytest.fixture +def event_log(): + return EventLog() + + +@pytest.fixture +def stop(): + return asyncio.Event() + + +@pytest.fixture +def file_lock(): + return asyncio.Lock() + + +# ── 跳过需要网络的测试(本地开发时) ── +def pytest_configure(config): + config.addinivalue_line("markers", "slow: 压力测试(较慢)") + config.addinivalue_line("markers", "needs_network: 需要外部网络") diff --git a/tests/fakes.py b/tests/fakes.py new file mode 100644 index 0000000..724798c --- /dev/null +++ b/tests/fakes.py @@ -0,0 +1,317 @@ +""" +Fake 服务 + 测试工具 + +FakeTurnstile — 可控 T 生成器:可注入延迟、失败、取消点 +FakeEmailService — 可控 Q 提供者:可注入延迟、超时、验证码 +FakeConsumer — 可控 C 消费器:可注入延迟、失败 +Conservation — slot 守恒断言器 +CancelPoint — 在指定 await 边界注入取消 +""" +import asyncio +import time +import dataclasses +from typing import Optional, Callable, Any + + +# ═══════════════════════════════════════════ +# Fake T 生成器 (替代 solve_one_turnstile) +# ═══════════════════════════════════════════ +class FakeTurnstile: + """可控 token 生成器。 + + config: + delay: 生成延迟(秒) + fail_rate: 0.0~1.0,随机失败概率 + token_seq: 预设 token 序列,None 则自动递增 + """ + + def __init__(self, delay=0.01, fail_rate=0.0, token_seq=None): + self.delay = delay + self._fail_rate = fail_rate + self._seq = token_seq or [] + self._idx = 0 + self._call_count = 0 + self._fail_count = 0 + self._cancel_at: Optional[int] = None # 第 N 次调用时注入取消 + + def set_cancel_at(self, n: int): + """第 n 次调用时注入 CancelledError。""" + self._cancel_at = n + + async def generate(self) -> Optional[str]: + """生成一个 token。返回 None 表示失败,raise CancelledError 表示取消。""" + self._call_count += 1 + + if self._cancel_at is not None and self._call_count == self._cancel_at: + raise asyncio.CancelledError(f'FakeTurnstile cancel at call {self._call_count}') + + await asyncio.sleep(self.delay) + + if self._fail_count < self._fail_rate * self._call_count: + return None + # 简单随机: fail_rate 作为概率 + import random + if random.random() < self._fail_rate: + self._fail_count += 1 + return None + + if self._idx < len(self._seq): + tok = self._seq[self._idx] + self._idx += 1 + return tok + self._idx += 1 + return f'fake_token_{self._idx}' + + @property + def call_count(self): + return self._call_count + + +# ═══════════════════════════════════════════ +# Fake 邮箱服务 (替代 create_email + poll_code) +# ═══════════════════════════════════════════ +@dataclasses.dataclass +class FakeEmailResult: + handle: str + email: str + password: str + code: Optional[str] = None # None = 轮询超时 + create_delay: float = 0.01 + poll_delay: float = 0.01 + + +class FakeEmailService: + """可控邮箱 + 验证码服务。 + + results: 预设结果队列,每次 create+poll 消费一个。 + """ + + def __init__(self, results=None): + self._results = list(results or []) + self._idx = 0 + self._create_count = 0 + self._poll_count = 0 + self._cancel_create_at: Optional[int] = None + self._cancel_poll_at: Optional[int] = None + + def set_cancel_create_at(self, n: int): + self._cancel_create_at = n + + def set_cancel_poll_at(self, n: int): + self._cancel_poll_at = n + + async def create_email(self): + """创建邮箱。返回 (handle, email, password)。""" + self._create_count += 1 + if self._cancel_create_at and self._create_count == self._cancel_create_at: + raise asyncio.CancelledError(f'FakeEmail.create cancel at {self._create_count}') + + if self._idx < len(self._results): + r = self._results[self._idx] + await asyncio.sleep(r.create_delay) + return r.handle, r.email, r.password + self._idx += 1 + n = self._idx + await asyncio.sleep(0.01) + return f'handle_{n}', f'user_{n}@test.com', f'pass_{n}' + + async def poll_code(self, handle: str, max_wait=90): + """轮询验证码。返回 code 或 None。""" + self._poll_count += 1 + if self._cancel_poll_at and self._poll_count == self._cancel_poll_at: + raise asyncio.CancelledError(f'FakeEmail.poll cancel at {self._poll_count}') + + # 找匹配的预设结果 + for r in self._results: + if r.handle == handle: + await asyncio.sleep(r.poll_delay) + return r.code + await asyncio.sleep(0.01) + return '000000' + + @property + def create_count(self): + return self._create_count + + @property + def poll_count(self): + return self._poll_count + + +# ═══════════════════════════════════════════ +# Fake 注册 API (替代 grpc_verify + server_action_register) +# ═══════════════════════════════════════════ +class FakeRegisterAPI: + """可控注册 API。""" + + def __init__(self, success_rate=1.0, delay=0.01): + self.success_rate = success_rate + self.delay = delay + self._register_count = 0 + self._cancel_at: Optional[int] = None + self._fail_at: Optional[int] = None + + def set_cancel_at(self, n: int): + self._cancel_at = n + + def set_fail_at(self, n: int): + self._fail_at = n + + async def register(self, email, password, code, token) -> Optional[str]: + """执行注册。返回 SSO 或 None。""" + self._register_count += 1 + if self._cancel_at and self._register_count == self._cancel_at: + raise asyncio.CancelledError(f'FakeRegister cancel at {self._register_count}') + if self._fail_at and self._register_count == self._fail_at: + return None + await asyncio.sleep(self.delay) + import random + if random.random() < self.success_rate: + return f'sso_{email}_{int(time.time())}' + return None + + @property + def register_count(self): + return self._register_count + + +# ═══════════════════════════════════════════ +# Slot 守恒断言器 +# ═══════════════════════════════════════════ +class Conservation: + """检查 T/Q slot 守恒不变量。 + + 守恒公式: + slot_total = free_slots + inventory_depth + pairleased_held + worker_held + """ + + def __init__(self, t_slot_sem, q_slot_sem, q_pending_sem, inventory): + self.t_slot_sem = t_slot_sem + self.q_slot_sem = q_slot_sem + self.q_pending_sem = q_pending_sem + self.inventory = inventory + # 追踪 Worker 持有的未发布资源 + self._worker_held_t = 0 + self._worker_held_q = 0 + + def worker_acquire_t(self): + self._worker_held_t += 1 + + def worker_release_t(self): + self._worker_held_t -= 1 + + def worker_acquire_q(self): + self._worker_held_q += 1 + + def worker_release_q(self): + self._worker_held_q -= 1 + + def check_t_conservation(self, t_slot_total: int, pairleased_t: int = 0): + """检查 T slot 守恒。 + + Args: + t_slot_total: T_Slot_Sem 初始容量 + pairleased_t: 当前 PairLease 持有的 T 数量 + """ + free = self.t_slot_sem._value + inv = self.inventory.t_depth + held = self._worker_held_t + total = free + inv + pairleased_t + held + assert total == t_slot_total, ( + f'T slot 守恒违反: free({free}) + inv({inv}) + pairleased({pairleased_t}) ' + f'+ worker_held({held}) = {total} != total({t_slot_total})' + ) + + def check_q_conservation(self, q_slot_total: int, pairleased_q: int = 0): + """检查 Q slot 守恒。""" + free = self.q_slot_sem._value + inv = self.inventory.q_depth + held = self._worker_held_q + total = free + inv + pairleased_q + held + assert total == q_slot_total, ( + f'Q slot 守恒违反: free({free}) + inv({inv}) + pairleased({pairleased_q}) ' + f'+ worker_held({held}) = {total} != total({q_slot_total})' + ) + + def check_pending_conservation(self, pending_total: int, in_flight: int = 0): + """检查 Q pending 守恒。""" + free = self.q_pending_sem._value + total = free + in_flight + assert total == pending_total, ( + f'Q pending 守恒违反: free({free}) + in_flight({in_flight}) = {total} ' + f'!= total({pending_total})' + ) + + def check_all(self, t_total, q_total, pend_total, pairleased_t=0, pairleased_q=0, in_flight=0): + """一次性检查所有守恒。""" + self.check_t_conservation(t_total, pairleased_t) + self.check_q_conservation(q_total, pairleased_q) + self.check_pending_conservation(pend_total, in_flight) + + +# ═══════════════════════════════════════════ +# 取消注入器 +# ═══════════════════════════════════════════ +class CancelInjector: + """在指定 await 边界前后注入取消。""" + + def __init__(self): + self._points: dict[str, asyncio.Event] = {} + self._cancel_task: Optional[asyncio.Task] = None + + def register(self, name: str): + """注册一个取消点。""" + self._points[name] = asyncio.Event() + + async def wait_at(self, name: str): + """Worker 在此挂起,等待外部触发取消。""" + if name in self._points: + await self._points[name].wait() + + def trigger(self, name: str, task: asyncio.Task): + """触发指定取消点,取消目标 task。""" + if name in self._points: + self._points[name].set() + task.cancel() + + async def cancel_after(self, delay: float, task: asyncio.Task): + """延迟后取消目标 task。""" + await asyncio.sleep(delay) + task.cancel() + + +# ═══════════════════════════════════════════ +# 事件记录器(用于性质测试) +# ═══════════════════════════════════════════ +@dataclasses.dataclass +class Event: + ts: float + kind: str # 't_produced', 't_admitted', 't_claimed', 'q_sent', 'q_returned', + # 'q_admitted', 'q_claimed', 'pair_claimed', 'pair_ok', 'pair_fail', + # 't_expired', 'q_expired', 't_discarded', 'q_discarded', + # 'sem_t_free', 'sem_q_free', 'sem_pend_free', 'inv_t', 'inv_q' + detail: Any = None + + +class EventLog: + """记录所有事件,用于事后分析。""" + + def __init__(self): + self.events: list[Event] = [] + self._lock = asyncio.Lock() + + async def record(self, kind: str, detail=None): + async with self._lock: + self.events.append(Event(ts=time.time(), kind=kind, detail=detail)) + + def count(self, kind: str) -> int: + return sum(1 for e in self.events if e.kind == kind) + + def filter(self, kind: str) -> list[Event]: + return [e for e in self.events if e.kind == kind] + + def dump(self) -> str: + lines = [] + for e in self.events: + lines.append(f'{e.ts:.3f} {e.kind} {e.detail or ""}') + return '\n'.join(lines) diff --git a/tests/requirements.txt b/tests/requirements.txt new file mode 100644 index 0000000..1354ca6 --- /dev/null +++ b/tests/requirements.txt @@ -0,0 +1,2 @@ +pytest>=8.0.0 +pytest-asyncio>=0.24.0 diff --git a/tests/test_admission_gate.py b/tests/test_admission_gate.py new file mode 100644 index 0000000..f99466e --- /dev/null +++ b/tests/test_admission_gate.py @@ -0,0 +1,82 @@ +import asyncio +import unittest + +from grok_register.core.admission import AdmissionGate + + +class FakeInventory: + def __init__(self): + self.t_depth = 0 + self.q_depth = 0 + + +class AdmissionGateTests(unittest.IsolatedAsyncioTestCase): + async def test_t_production_blocks_at_high_water_and_wakes_at_low_water(self): + inventory = FakeInventory() + gate = AdmissionGate(inventory, t_low=1, t_high=3, q_low=1, q_high=3) + + inventory.t_depth = 3 + blocked = asyncio.create_task(gate.acquire_t_production()) + await asyncio.sleep(0.02) + self.assertFalse(blocked.done()) + + inventory.t_depth = 1 + await gate.notify_changed() + lease = await asyncio.wait_for(blocked, timeout=1) + + self.assertEqual(gate.t_in_progress, 1) + await lease.release() + self.assertEqual(gate.t_in_progress, 0) + + async def test_q_batch_reserves_only_current_demand(self): + inventory = FakeInventory() + gate = AdmissionGate(inventory, t_low=1, t_high=3, q_low=2, q_high=5) + + lease = await gate.acquire_q_batch(max_batch=4) + + self.assertEqual(lease.count, 4) + self.assertEqual(gate.q_inflight, 4) + + second = await gate.acquire_q_batch(max_batch=4) + self.assertEqual(second.count, 1) + self.assertEqual(gate.q_inflight, 5) + + blocked = asyncio.create_task(gate.acquire_q_batch(max_batch=4)) + await asyncio.sleep(0.02) + self.assertFalse(blocked.done()) + + await lease.release_all() + await second.release_all() + inventory.q_depth = 1 + await gate.notify_changed() + third = await asyncio.wait_for(blocked, timeout=1) + self.assertEqual(third.count, 4) + await third.release_all() + + async def test_q_batch_waits_for_low_water_after_reaching_high_water(self): + inventory = FakeInventory() + gate = AdmissionGate(inventory, t_low=1, t_high=3, q_low=2, q_high=5) + + lease = await gate.acquire_q_batch(max_batch=5) + self.assertEqual(lease.count, 5) + self.assertEqual(gate.q_inflight, 5) + + await lease.release_one() + self.assertEqual(gate.q_inflight, 4) + + blocked = asyncio.create_task(gate.acquire_q_batch(max_batch=5)) + await asyncio.sleep(0.02) + self.assertFalse(blocked.done()) + + await lease.release_one() + await lease.release_one() + await lease.release_one() + + resumed = await asyncio.wait_for(blocked, timeout=1) + self.assertEqual(resumed.count, 4) + await lease.release_all() + await resumed.release_all() + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_cancel.py b/tests/test_cancel.py new file mode 100644 index 0000000..9c9ad60 --- /dev/null +++ b/tests/test_cancel.py @@ -0,0 +1,475 @@ +""" +Layer 2: 取消注入测试 + +在每个 await 边界前后强制取消 Worker,验证: + - S 取消后 Physical_Sem 和 T_Slot_Sem 不泄漏 + - P 取消后 Physical_Sem、Q_Pending_Sem、Q_Slot_Sem 不泄漏 + - C 取消后 Physical_Sem 不泄漏,已 claim 的 T/Q 被核销 + - 库存守恒: 取消不导致 slot 无界堆积或耗尽 +""" +import asyncio +import time +import pytest + +from grok_register.core.envelope import ResourceEnvelope +from grok_register.core.inventory import Inventory +from grok_register.core.observer import Metrics +from tests.fakes import FakeTurnstile, FakeEmailService, FakeRegisterAPI, Conservation + + +# ═══════════════════════════════════════════ +# S_Worker 取消测试 +# ═══════════════════════════════════════════ + +async def _s_worker_impl(browser, inventory, physical_sem, t_slot_sem, metrics, stop): + """S_Worker 简化实现,用于测试。""" + while not stop.is_set(): + await physical_sem.acquire() + try: + token = await browser.generate() + finally: + physical_sem.release() + + if token is None: + metrics.t_discarded += 1 + await asyncio.sleep(0.05) + continue + + metrics.t_produced += 1 + now = time.time() + try: + t_env = await ResourceEnvelope.create_with_slot( + 'T', token, t_slot_sem, expires_at=now + 300 + ) + except asyncio.CancelledError: + metrics.t_discarded += 1 + raise + + try: + await inventory.put_t(t_env) + except Exception: + t_env.discard(reason='put_failed') + metrics.t_discarded += 1 + await asyncio.sleep(0.02) + + +class TestSCancel: + """S_Worker 取消语义测试。""" + + @pytest.mark.asyncio + async def test_cancel_during_generate_no_leak(self): + """S 在 generate() 期间被取消,sem 不泄漏。""" + ft = FakeTurnstile(delay=0.5) # 慢生成 + inv = Inventory() + phys = asyncio.Semaphore(2) + t_slot = asyncio.Semaphore(4) + metrics = Metrics() + stop = asyncio.Event() + + task = asyncio.create_task( + _s_worker_impl(ft, inv, phys, t_slot, metrics, stop) + ) + await asyncio.sleep(0.05) # 等它进入 generate + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + # Physical_Sem 应完全释放 + assert phys._value == 2, f'Physical_Sem 泄漏: {phys._value}' + # T_Slot_Sem 不变(没生成成功) + assert t_slot._value == 4 + + @pytest.mark.asyncio + async def test_cancel_during_put_t_no_leak(self): + """S 在 put_t 期间被取消(极端: put 内部取消),envelope 被 discard。""" + inv = Inventory() + phys = asyncio.Semaphore(2) + t_slot = asyncio.Semaphore(4) + metrics = Metrics() + stop = asyncio.Event() + + # 快速生成 + ft = FakeTurnstile(delay=0.001) + + task = asyncio.create_task( + _s_worker_impl(ft, inv, phys, t_slot, metrics, stop) + ) + # 在 put_t 之前的极短窗口取消 + await asyncio.sleep(0.002) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert phys._value == 2 + # t_slot 可能是 4 或 3,但不应是负数 + assert t_slot._value >= 0 + assert t_slot._value <= 4 + + @pytest.mark.asyncio + async def test_s_repeated_cancel_no_accumulation(self): + """反复取消 S 多次,sem 计数不累积。""" + ft = FakeTurnstile(delay=0.01) + inv = Inventory() + phys = asyncio.Semaphore(2) + t_slot = asyncio.Semaphore(4) + metrics = Metrics() + stop = asyncio.Event() + + for _ in range(20): + task = asyncio.create_task( + _s_worker_impl(ft, inv, phys, t_slot, metrics, stop) + ) + await asyncio.sleep(0.02) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert phys._value == 2, f'Physical_Sem 泄漏: {phys._value}' + # 守恒: free + inventory == total (成功入库的 envelope 占着 slot) + assert t_slot._value + inv.t_depth == 4, f'T_Slot 守恒违反: free={t_slot._value} inv={inv.t_depth}' + + +# ═══════════════════════════════════════════ +# P_Worker 取消测试 +# ═══════════════════════════════════════════ + +async def _p_worker_impl(email_svc, inventory, physical_sem, q_pending_sem, + q_slot_sem, metrics, stop): + """P_Worker 简化实现,用于测试。""" + loop = asyncio.get_event_loop() + while not stop.is_set(): + await q_pending_sem.acquire() + _pending_released = False + try: + await physical_sem.acquire() + try: + handle, email, password = await email_svc.create_email() + sent = True # fake 总是成功 + finally: + physical_sem.release() + + if not sent: + q_pending_sem.release() + _pending_released = True + continue + + metrics.q_sent += 1 + + try: + code = await asyncio.wait_for( + email_svc.poll_code(handle), timeout=95 + ) + except asyncio.TimeoutError: + code = None + + if code is None: + metrics.q_discarded += 1 + q_pending_sem.release() + _pending_released = True + continue + + metrics.q_returned += 1 + now = time.time() + q_env = await ResourceEnvelope.create_with_slot( + 'Q', {'email': email, 'password': password, 'code': code}, + q_slot_sem, expires_at=now + 120, + ) + + try: + await inventory.put_q(q_env) + except Exception: + q_env.discard(reason='put_failed') + metrics.q_discarded += 1 + + except asyncio.CancelledError: + if not _pending_released: + q_pending_sem.release() + _pending_released = True + raise + except Exception: + metrics.q_discarded += 1 + finally: + if not _pending_released: + q_pending_sem.release() + + await asyncio.sleep(0.02) + + +class TestPCancel: + """P_Worker 取消语义测试。""" + + @pytest.mark.asyncio + async def test_cancel_during_create_email(self): + """P 在 create_email 期间被取消,所有 sem 正确释放。""" + es = FakeEmailService() + es.set_cancel_create_at(1) + inv = Inventory() + phys = asyncio.Semaphore(2) + q_pend = asyncio.Semaphore(4) + q_slot = asyncio.Semaphore(4) + metrics = Metrics() + stop = asyncio.Event() + + task = asyncio.create_task( + _p_worker_impl(es, inv, phys, q_pend, q_slot, metrics, stop) + ) + await asyncio.sleep(0.1) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert phys._value == 2, f'Physical_Sem 泄漏: {phys._value}' + assert q_pend._value == 4, f'Q_Pending 泄漏: {q_pend._value}' + assert q_slot._value == 4 + + @pytest.mark.asyncio + async def test_cancel_during_poll_code(self): + """P 在 poll_code 期间被取消,不持有 Physical_Sem。""" + es = FakeEmailService() + es.set_cancel_poll_at(1) + inv = Inventory() + phys = asyncio.Semaphore(2) + q_pend = asyncio.Semaphore(4) + q_slot = asyncio.Semaphore(4) + metrics = Metrics() + stop = asyncio.Event() + + task = asyncio.create_task( + _p_worker_impl(es, inv, phys, q_pend, q_slot, metrics, stop) + ) + await asyncio.sleep(0.1) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + # P 等 Q 时不持有 Physical_Sem + assert phys._value == 2, f'Physical_Sem 应完全释放: {phys._value}' + assert q_pend._value == 4, f'Q_Pending 泄漏: {q_pend._value}' + + @pytest.mark.asyncio + async def test_p_repeated_cancel_no_accumulation(self): + """反复取消 P 多次,sem 计数不累积。""" + es = FakeEmailService() + inv = Inventory() + phys = asyncio.Semaphore(2) + q_pend = asyncio.Semaphore(4) + q_slot = asyncio.Semaphore(4) + metrics = Metrics() + stop = asyncio.Event() + + for _ in range(20): + task = asyncio.create_task( + _p_worker_impl(es, inv, phys, q_pend, q_slot, metrics, stop) + ) + await asyncio.sleep(0.03) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert phys._value == 2, f'Physical_Sem 泄漏: {phys._value}' + assert q_pend._value == 4, f'Q_Pending 泄漏: {q_pend._value}' + # 守恒: free + inventory == total + assert q_slot._value + inv.q_depth == 4, f'Q_Slot 守恒违反: free={q_slot._value} inv={inv.q_depth}' + + +# ═══════════════════════════════════════════ +# C_Worker 取消测试 +# ═══════════════════════════════════════════ + +async def _c_worker_impl(register_api, browser_unused, inventory, physical_sem, + metrics, stop, file_lock): + """C_Worker 简化实现,用于测试。""" + while not stop.is_set(): + try: + async with inventory.claim_pair() as pair: + t_val = pair.t.value + q_val = pair.q.value + await physical_sem.acquire() + try: + sso = await register_api.register( + q_val['email'], q_val['password'], q_val['code'], t_val + ) + if sso: + metrics.pair_consumed_ok += 1 + metrics.success_count += 1 + else: + metrics.pair_consumed_fail += 1 + finally: + physical_sem.release() + except asyncio.CancelledError: + raise + except Exception: + metrics.pair_consumed_fail += 1 + await asyncio.sleep(0.02) + + +class TestCCancel: + """C_Worker 取消语义测试。""" + + @pytest.mark.asyncio + async def test_cancel_during_register_pair_settled(self): + """C 在 register 期间被取消,pair 已 claim,必须核销。""" + fra = FakeRegisterAPI(delay=0.5) + inv = Inventory() + phys = asyncio.Semaphore(2) + t_slot = asyncio.Semaphore(4) + q_slot = asyncio.Semaphore(4) + metrics = Metrics() + stop = asyncio.Event() + fl = asyncio.Lock() + + # 用带 metrics 的 inventory + inv_with_metrics = Inventory(metrics=Metrics()) + t_env = await ResourceEnvelope.create_with_slot('T', 'tok', t_slot) + q_env = await ResourceEnvelope.create_with_slot('Q', {'email': 'a@b.com', 'password': 'p', 'code': '123'}, q_slot) + await inv_with_metrics.put_t(t_env) + await inv_with_metrics.put_q(q_env) + + task = asyncio.create_task( + _c_worker_impl(fra, None, inv_with_metrics, phys, metrics, stop, fl) + ) + await asyncio.sleep(0.1) # 等它 claim 进入 register + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + # pair 已核销: slot 应释放 + assert t_slot._value == 4, f'T slot 未核销: {t_slot._value}' + assert q_slot._value == 4, f'Q slot 未核销: {q_slot._value}' + assert phys._value == 2, f'Physical_Sem 泄漏: {phys._value}' + + @pytest.mark.asyncio + async def test_cancel_during_claim_wait_no_pop(self): + """C 在 claim_pair 等待时取消,不弹出资源。""" + inv = Inventory() + phys = asyncio.Semaphore(2) + t_slot = asyncio.Semaphore(4) + q_slot = asyncio.Semaphore(4) + metrics = Metrics() + stop = asyncio.Event() + fl = asyncio.Lock() + + fra = FakeRegisterAPI() + + # 放一个 T 但不放 Q,让 C 等待 + t_env = await ResourceEnvelope.create_with_slot('T', 'tok', t_slot) + await inv.put_t(t_env) + + task = asyncio.create_task( + _c_worker_impl(fra, None, inv, phys, metrics, stop, fl) + ) + await asyncio.sleep(0.1) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + # T 仍在 inventory 中 + assert inv.t_depth == 1, 'claim 等待时取消不应弹出 T' + assert phys._value == 2 + + @pytest.mark.asyncio + async def test_c_repeated_cancel_no_accumulation(self): + """反复取消 C 多次,sem 计数不累积。""" + fra = FakeRegisterAPI(delay=0.01) + inv = Inventory() + phys = asyncio.Semaphore(2) + t_slot = asyncio.Semaphore(10) + q_slot = asyncio.Semaphore(10) + metrics = Metrics() + stop = asyncio.Event() + fl = asyncio.Lock() + + for _ in range(10): + # 每次放一对 + t_env = await ResourceEnvelope.create_with_slot('T', f'tok_{_}', t_slot) + q_env = await ResourceEnvelope.create_with_slot('Q', {'email': f'e@{_}', 'password': 'p', 'code': '123'}, q_slot) + await inv.put_t(t_env) + await inv.put_q(q_env) + + task = asyncio.create_task( + _c_worker_impl(fra, None, inv, phys, metrics, stop, fl) + ) + await asyncio.sleep(0.03) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert phys._value == 2, f'Physical_Sem 泄漏: {phys._value}' + # 所有 slot 应已核销 + assert t_slot._value == 10, f'T_Slot 泄漏: {t_slot._value}' + assert q_slot._value == 10, f'Q_Slot 泄漏: {q_slot._value}' + + +# ═══════════════════════════════════════════ +# 端到端取消: S+P+C 同时运行时取消 +# ═══════════════════════════════════════════ +class TestE2ECancel: + """端到端取消测试: S/P/C 同时运行时取消,验证全局守恒。""" + + @pytest.mark.asyncio + async def test_cancel_all_workers_sem_reset(self): + """同时启动 S/P/C,然后全部取消,所有 sem 恢复初始值。""" + ft = FakeTurnstile(delay=0.05) + es = FakeEmailService() + fra = FakeRegisterAPI(delay=0.05) + + inv = Inventory() + phys = asyncio.Semaphore(4) + t_slot = asyncio.Semaphore(8) + q_slot = asyncio.Semaphore(8) + q_pend = asyncio.Semaphore(12) + metrics = Metrics() + stop = asyncio.Event() + fl = asyncio.Lock() + + tasks = [] + for _ in range(3): + tasks.append(asyncio.create_task( + _s_worker_impl(ft, inv, phys, t_slot, metrics, stop) + )) + tasks.append(asyncio.create_task( + _p_worker_impl(es, inv, phys, q_pend, q_slot, metrics, stop) + )) + tasks.append(asyncio.create_task( + _c_worker_impl(fra, None, inv, phys, metrics, stop, fl) + )) + + await asyncio.sleep(0.2) + stop.set() + for t in tasks: + t.cancel() + for t in tasks: + try: + await t + except asyncio.CancelledError: + pass + + # 全部停止后,守恒检查: free + inventory == total + assert phys._value >= 0 and phys._value <= 4 + assert t_slot._value + inv.t_depth == 8, f'T 守恒违反: free={t_slot._value} inv={inv.t_depth}' + assert q_slot._value + inv.q_depth == 8, f'Q 守恒违反: free={q_slot._value} inv={inv.q_depth}' + assert q_pend._value >= 0 and q_pend._value <= 12 + + # 守恒检查 + cons = Conservation(t_slot, q_slot, q_pend, inv) + cons.check_t_conservation(8) + cons.check_q_conservation(8) + cons.check_pending_conservation(12) diff --git a/tests/test_core.py b/tests/test_core.py new file mode 100644 index 0000000..5a97cac --- /dev/null +++ b/tests/test_core.py @@ -0,0 +1,491 @@ +""" +Layer 1: 单元测试 + +测 ResourceEnvelope、PairLease、Inventory 的局部语义: + - release_slot_once() 幂等 + - create_with_slot() slot 获取失败不创建 envelope + - is_expired() 正确判断 + - put_t/put_q 成功前所有权属于 Worker,成功后属于 Inventory + - claim_pair 等待时取消不弹资源 + - claim_pair 成功后 pair 属于 PairLease + - PairLease 退出时核销 T/Q + - 过期资源不会交付给 C + - lazy cleanup 只在事件触发时发生 + - discard() 幂等释放 +""" +import asyncio +import time +import pytest + +from grok_register.core.envelope import ResourceEnvelope +from grok_register.core.inventory import Inventory, PairLease + + +# ═══════════════════════════════════════════ +# ResourceEnvelope 测试 +# ═══════════════════════════════════════════ +class TestResourceEnvelope: + """ResourceEnvelope 语义测试。""" + + @pytest.mark.asyncio + async def test_create_with_slot_acquires_sem(self): + """create_with_slot 成功后 semaphore 值减 1。""" + sem = asyncio.Semaphore(3) + env = await ResourceEnvelope.create_with_slot('T', 'tok1', sem) + assert sem._value == 2 + assert env.value == 'tok1' + assert env.kind == 'T' + assert not env.released + + @pytest.mark.asyncio + async def test_create_with_slot_fails_no_envelope(self): + """slot 获取失败(取消)时不创建 envelope,sem 不变。""" + sem = asyncio.Semaphore(0) + with pytest.raises(asyncio.CancelledError): + # 手动取消 + task = asyncio.current_task() + task.cancel() + try: + await ResourceEnvelope.create_with_slot('T', 'tok1', sem) + except asyncio.CancelledError: + # sem 不应改变(因为 acquire 被取消了) + raise + # sem 仍为 0 + assert sem._value == 0 + + @pytest.mark.asyncio + async def test_release_slot_once_idempotent(self): + """release_slot_once() 幂等:重复调用只释放一次。""" + sem = asyncio.Semaphore(3) + env = await ResourceEnvelope.create_with_slot('T', 'tok1', sem) + assert sem._value == 2 + + env.release_slot_once(reason='test1') + assert sem._value == 3 + assert env.released + + # 第二次调用不应再次释放 + env.release_slot_once(reason='test2') + assert sem._value == 3 + + @pytest.mark.asyncio + async def test_discard_releases_slot(self): + """discard() 释放 slot。""" + sem = asyncio.Semaphore(2) + env = await ResourceEnvelope.create_with_slot('Q', {'code': '123'}, sem) + assert sem._value == 1 + env.discard(reason='test') + assert sem._value == 2 + assert env.released + + @pytest.mark.asyncio + async def test_discard_idempotent(self): + """discard() 幂等。""" + sem = asyncio.Semaphore(2) + env = await ResourceEnvelope.create_with_slot('Q', 'val', sem) + env.discard(reason='first') + env.discard(reason='second') + assert sem._value == 2 + + @pytest.mark.asyncio + async def test_is_expired_with_expires_at(self): + """有 expires_at 时正确判断过期。""" + sem = asyncio.Semaphore(2) + now = time.time() + env = await ResourceEnvelope.create_with_slot('T', 'tok', sem, expires_at=now - 1) + assert env.is_expired(now=now) + + env2 = await ResourceEnvelope.create_with_slot('T', 'tok2', sem, expires_at=now + 100) + assert not env2.is_expired(now=now) + + @pytest.mark.asyncio + async def test_is_expired_without_expires_at(self): + """没有 expires_at 时永远不过期。""" + sem = asyncio.Semaphore(2) + env = await ResourceEnvelope.create_with_slot('T', 'tok', sem) + assert not env.is_expired(now=time.time() + 99999) + + @pytest.mark.asyncio + async def test_create_with_slot_meta(self): + """meta 字段正确传递。""" + sem = asyncio.Semaphore(1) + env = await ResourceEnvelope.create_with_slot('T', 'tok', sem, meta={'source': 'test'}) + assert env.meta == {'source': 'test'} + + +# ═══════════════════════════════════════════ +# Inventory 测试 +# ═══════════════════════════════════════════ +class TestInventory: + """Inventory 语义测试。""" + + @pytest.mark.asyncio + async def test_put_t_then_claim(self): + """put_t 后 claim_pair 能拿到 pair。""" + inv = Inventory() + t_sem = asyncio.Semaphore(2) + q_sem = asyncio.Semaphore(2) + + t_env = await ResourceEnvelope.create_with_slot('T', 'token1', t_sem) + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': 'abc'}, q_sem) + + await inv.put_t(t_env) + await inv.put_q(q_env) + + assert inv.t_depth == 1 + assert inv.q_depth == 1 + + # claim + async with inv.claim_pair() as pair: + assert pair.t.value == 'token1' + assert pair.q.value == {'code': 'abc'} + # claim 后 inventory 应为空 + assert inv.t_depth == 0 + assert inv.q_depth == 0 + + # PairLease 退出后 slot 应被释放 + assert t_sem._value == 2 + assert q_sem._value == 2 + + @pytest.mark.asyncio + async def test_claim_waits_for_both(self): + """只有 T 没有 Q 时,claim_pair 应等待。""" + inv = Inventory() + t_sem = asyncio.Semaphore(2) + q_sem = asyncio.Semaphore(2) + + t_env = await ResourceEnvelope.create_with_slot('T', 'tok', t_sem) + await inv.put_t(t_env) + + claimed = False + + async def claimer(): + nonlocal claimed + async with inv.claim_pair() as pair: + claimed = True + + task = asyncio.create_task(claimer()) + await asyncio.sleep(0.1) + assert not claimed, 'claim 不应在缺少 Q 时成功' + + # 现在放入 Q + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': 'x'}, q_sem) + await inv.put_q(q_env) + + await asyncio.wait_for(task, timeout=2) + assert claimed + + @pytest.mark.asyncio + async def test_claim_cancel_no_pop(self): + """claim_pair 等待时取消,不应弹出任何资源。""" + inv = Inventory() + t_sem = asyncio.Semaphore(2) + + t_env = await ResourceEnvelope.create_with_slot('T', 'tok', t_sem) + await inv.put_t(t_env) + assert inv.t_depth == 1 + + async def claimer(): + async with inv.claim_pair() as pair: + pass # 不应到达 + + task = asyncio.create_task(claimer()) + await asyncio.sleep(0.1) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + # T 仍在 inventory 中 + assert inv.t_depth == 1 + + @pytest.mark.asyncio + async def test_pair_lease_settles_on_success(self): + """PairLease 成功退出时核销 T/Q 并释放 slot。""" + inv = Inventory() + t_sem = asyncio.Semaphore(2) + q_sem = asyncio.Semaphore(2) + + t_env = await ResourceEnvelope.create_with_slot('T', 'tok', t_sem) + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': 'x'}, q_sem) + await inv.put_t(t_env) + await inv.put_q(q_env) + + async with inv.claim_pair() as pair: + # slot 还在 pair 中 + assert t_sem._value == 1 + assert q_sem._value == 1 + + # 退出后 slot 释放 + assert t_sem._value == 2 + assert q_sem._value == 2 + + @pytest.mark.asyncio + async def test_pair_lease_settles_on_exception(self): + """PairLease 异常退出时也核销 T/Q。""" + inv = Inventory() + t_sem = asyncio.Semaphore(2) + q_sem = asyncio.Semaphore(2) + + t_env = await ResourceEnvelope.create_with_slot('T', 'tok', t_sem) + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': 'x'}, q_sem) + await inv.put_t(t_env) + await inv.put_q(q_env) + + with pytest.raises(ValueError): + async with inv.claim_pair() as pair: + raise ValueError('boom') + + # slot 仍被核销 + assert t_sem._value == 2 + assert q_sem._value == 2 + + @pytest.mark.asyncio + async def test_pair_lease_settles_on_cancel(self): + """PairLease 取消退出时也核销 T/Q。""" + inv = Inventory() + t_sem = asyncio.Semaphore(2) + q_sem = asyncio.Semaphore(2) + + t_env = await ResourceEnvelope.create_with_slot('T', 'tok', t_sem) + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': 'x'}, q_sem) + await inv.put_t(t_env) + await inv.put_q(q_env) + + async def consumer(): + async with inv.claim_pair() as pair: + # 在 pair 内部被取消 + await asyncio.sleep(999) + + task = asyncio.create_task(consumer()) + await asyncio.sleep(0.1) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + # slot 仍被核销 + assert t_sem._value == 2 + assert q_sem._value == 2 + + @pytest.mark.asyncio + async def test_expired_t_not_delivered(self): + """过期的 T 不会交付给 C。""" + inv = Inventory() + t_sem = asyncio.Semaphore(2) + q_sem = asyncio.Semaphore(2) + + now = time.time() + # T 已过期 + t_env = await ResourceEnvelope.create_with_slot('T', 'expired_tok', t_sem, expires_at=now - 10) + # Q 有效 + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': 'x'}, q_sem, expires_at=now + 100) + + await inv.put_t(t_env) + await inv.put_q(q_env) + + # claim 应该把过期的 T 清理掉,然后等待新的 T + claimed = False + + async def claimer(): + nonlocal claimed + async with inv.claim_pair() as pair: + claimed = True + + task = asyncio.create_task(claimer()) + await asyncio.sleep(0.2) + assert not claimed, '过期 T 不应交付' + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + # 过期 T 被清理,slot 被释放 + assert t_sem._value == 2 + # Q 仍在 inventory 中 + assert inv.q_depth == 1 + + @pytest.mark.asyncio + async def test_expired_q_not_delivered(self): + """过期的 Q 不会交付给 C。""" + inv = Inventory() + t_sem = asyncio.Semaphore(2) + q_sem = asyncio.Semaphore(2) + + now = time.time() + t_env = await ResourceEnvelope.create_with_slot('T', 'tok', t_sem, expires_at=now + 100) + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': 'x'}, q_sem, expires_at=now - 10) + + await inv.put_t(t_env) + await inv.put_q(q_env) + + claimed = False + + async def claimer(): + nonlocal claimed + async with inv.claim_pair() as pair: + claimed = True + + task = asyncio.create_task(claimer()) + await asyncio.sleep(0.2) + assert not claimed, '过期 Q 不应交付' + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert q_sem._value == 2 + assert inv.t_depth == 1 + + @pytest.mark.asyncio + async def test_lazy_cleanup_on_put(self): + """lazy cleanup: put 时清理已过期的库存。""" + inv = Inventory() + t_sem = asyncio.Semaphore(4) + + now = time.time() + # 放入 3 个过期的 T + for i in range(3): + env = await ResourceEnvelope.create_with_slot( + 'T', f'expired_{i}', t_sem, expires_at=now - 10 + ) + await inv.put_t(env) + + # 由于 put 时 lazy cleanup 已清理过期项,库存应为空 + assert inv.t_depth == 0 + # slot 应已释放 + assert t_sem._value == 4 + + @pytest.mark.asyncio + async def test_lazy_cleanup_silent_until_next_event(self): + """lazy cleanup 不在静默期主动释放,下一次 put 才触发清理。""" + inv = Inventory() + t_sem = asyncio.Semaphore(4) + + env = await ResourceEnvelope.create_with_slot( + 'T', 'will_expire', t_sem, expires_at=time.time() + 0.05 + ) + await inv.put_t(env) + + await asyncio.sleep(0.1) + assert inv.t_depth == 1 + assert t_sem._value == 3 + + fresh = await ResourceEnvelope.create_with_slot( + 'T', 'fresh', t_sem, expires_at=time.time() + 100 + ) + await inv.put_t(fresh) + + assert inv.t_depth == 1 + assert t_sem._value == 3 + + @pytest.mark.asyncio + async def test_lazy_cleanup_on_claim(self): + """lazy cleanup: claim_pair 时清理过期库存。""" + inv = Inventory() + t_sem = asyncio.Semaphore(4) + q_sem = asyncio.Semaphore(4) + + now = time.time() + # 直接注入过期项到 inventory 内部(绕过 put 的清理) + env = await ResourceEnvelope.create_with_slot('T', 'expired', t_sem, expires_at=now - 10) + inv._t_buf.append(env) + + # 有效的 Q + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': 'x'}, q_sem, expires_at=now + 100) + await inv.put_q(q_env) + + # claim 时应清理过期的 T + claimed = False + async def claimer(): + nonlocal claimed + async with inv.claim_pair() as pair: + claimed = True + + task = asyncio.create_task(claimer()) + await asyncio.sleep(0.2) + assert not claimed # T 被清了,没有 pair + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert inv.t_depth == 0 + assert t_sem._value == 4 # slot 已释放 + + @pytest.mark.asyncio + async def test_fifo_order(self): + """Inventory 按 FIFO 顺序配对。""" + inv = Inventory() + t_sem = asyncio.Semaphore(10) + q_sem = asyncio.Semaphore(10) + + for i in range(3): + t_env = await ResourceEnvelope.create_with_slot('T', f'tok_{i}', t_sem) + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': f'c_{i}'}, q_sem) + await inv.put_t(t_env) + await inv.put_q(q_env) + + pairs = [] + for _ in range(3): + async with inv.claim_pair() as pair: + pairs.append((pair.t.value, pair.q.value['code'])) + + assert pairs == [('tok_0', 'c_0'), ('tok_1', 'c_1'), ('tok_2', 'c_2')] + + @pytest.mark.asyncio + async def test_multiple_concurrent_claims(self): + """多个 C 同时 claim,不应重复消费同一个 T/Q。""" + inv = Inventory() + t_sem = asyncio.Semaphore(10) + q_sem = asyncio.Semaphore(10) + + n = 5 + for i in range(n): + t_env = await ResourceEnvelope.create_with_slot('T', f'tok_{i}', t_sem) + q_env = await ResourceEnvelope.create_with_slot('Q', {'code': f'c_{i}'}, q_sem) + await inv.put_t(t_env) + await inv.put_q(q_env) + + results = [] + + async def claimer(idx): + async with inv.claim_pair() as pair: + results.append((idx, pair.t.value)) + + tasks = [asyncio.create_task(claimer(i)) for i in range(n)] + await asyncio.gather(*tasks) + + # 每个 T 只被消费一次 + t_values = [r[1] for r in results] + assert len(set(t_values)) == n, f'重复消费: {t_values}' + + @pytest.mark.asyncio + async def test_metrics_admitted(self): + """put 成功时 metrics.t_admitted / q_admitted 递增。""" + from grok_register.core.observer import Metrics + m = Metrics() + inv = Inventory(metrics=m) + t_sem = asyncio.Semaphore(2) + + env = await ResourceEnvelope.create_with_slot('T', 'tok', t_sem) + await inv.put_t(env) + assert m.t_admitted == 1 + + @pytest.mark.asyncio + async def test_metrics_expired(self): + """过期资源 put 时 metrics.t_expired 递增。""" + from grok_register.core.observer import Metrics + m = Metrics() + inv = Inventory(metrics=m) + t_sem = asyncio.Semaphore(2) + + env = await ResourceEnvelope.create_with_slot('T', 'tok', t_sem, expires_at=time.time() - 10) + await inv.put_t(env) + assert m.t_expired == 1 + assert m.t_admitted == 0 diff --git a/tests/test_inventory_unittest.py b/tests/test_inventory_unittest.py new file mode 100644 index 0000000..92e85c4 --- /dev/null +++ b/tests/test_inventory_unittest.py @@ -0,0 +1,58 @@ +import asyncio +import time +import unittest + +from grok_register.core.envelope import ResourceEnvelope +from grok_register.core.inventory import Inventory +from grok_register.core.observer import Metrics + + +class InventoryUnitTests(unittest.IsolatedAsyncioTestCase): + async def test_create_with_slot_releases_slot_if_constructor_fails(self): + class BrokenEnvelope(ResourceEnvelope): + def __init__(self, *args, **kwargs): + raise RuntimeError("boom") + + sem = asyncio.Semaphore(1) + + with self.assertRaises(RuntimeError): + await BrokenEnvelope.create_with_slot("T", "tok", sem) + + self.assertEqual(sem._value, 1) + + async def test_lazy_cleanup_scans_past_fresh_head(self): + metrics = Metrics() + inventory = Inventory(metrics=metrics) + t_sem = asyncio.Semaphore(3) + q_sem = asyncio.Semaphore(1) + now = time.time() + + fresh_head = await ResourceEnvelope.create_with_slot( + "T", "fresh_head", t_sem, expires_at=now + 100 + ) + stale_tail = await ResourceEnvelope.create_with_slot( + "T", "stale_tail", t_sem, expires_at=now + 100 + ) + fresh_tail = await ResourceEnvelope.create_with_slot( + "T", "fresh_tail", t_sem, expires_at=now + 100 + ) + + await inventory.put_t(fresh_head) + await inventory.put_t(stale_tail) + await inventory.put_t(fresh_tail) + self.assertEqual(t_sem._value, 0) + + stale_tail.expires_at = time.time() - 10 + + q_env = await ResourceEnvelope.create_with_slot( + "Q", "q", q_sem, expires_at=now + 100 + ) + await inventory.put_q(q_env) + + self.assertEqual(t_sem._value, 1) + self.assertEqual(metrics.t_expired, 1) + self.assertEqual(inventory.t_depth, 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_property.py b/tests/test_property.py new file mode 100644 index 0000000..de7c2cc --- /dev/null +++ b/tests/test_property.py @@ -0,0 +1,482 @@ +""" +Layer 3: 性质测试 (Property Tests) + +随机生成 S/P/C 成功、失败、取消、超时、过期事件,持续检查: + - slot 守恒 (T/Q 总量恒定) + - 无重复 claim (同一个 T/Q 不被两个 PairLease 拿到) + - 无单边持有等待 (C 不先拿 T 再等 Q) + - 过期资源不交付 + - 取消后不泄漏 +""" +import asyncio +import time +import random +import pytest + +from grok_register.core.envelope import ResourceEnvelope +from grok_register.core.inventory import Inventory +from grok_register.core.observer import Metrics +from tests.fakes import FakeTurnstile, FakeEmailService, FakeRegisterAPI, Conservation, EventLog + + +# ═══════════════════════════════════════════ +# 随机性质测试引擎 +# ═══════════════════════════════════════════ + +class PropertyEngine: + """驱动 S/P/C 随机行为,持续检查不变量。""" + + def __init__(self, t_slot_cap=6, q_slot_cap=6, q_pending_cap=8, phys_cap=4, + t_max_age=2.0, q_max_age=2.0, seed=None): + self.rng = random.Random(seed) + self.t_slot_sem = asyncio.Semaphore(t_slot_cap) + self.q_slot_sem = asyncio.Semaphore(q_slot_cap) + self.q_pending_sem = asyncio.Semaphore(q_pending_cap) + self.phys_sem = asyncio.Semaphore(phys_cap) + self.t_slot_cap = t_slot_cap + self.q_slot_cap = q_slot_cap + self.q_pending_cap = q_pending_cap + self.t_max_age = t_max_age + self.q_max_age = q_max_age + + self.metrics = Metrics() + self.inventory = Inventory(metrics=self.metrics) + self.conservation = Conservation( + self.t_slot_sem, self.q_slot_sem, self.q_pending_sem, self.inventory + ) + self.events = EventLog() + + self._claimed_tokens = set() # 检测重复 claim + self._pairleased_t = 0 # 当前 PairLease 持有的 T + self._pairleased_q = 0 # 当前 PairLease 持有的 Q + + def check_invariants_soft(self): + """运行时软检查: 非负、不超限。 + + 注意: Worker 运行期间 slot 可能在 create_with_slot 和 put_t 之间的 + envelope 中(worker 持有但无法从外部追踪)。严格守恒检查用 + check_invariants_final() 在所有 Worker 停止后执行。 + """ + assert self.t_slot_sem._value >= 0, 'T slot 为负' + assert self.q_slot_sem._value >= 0, 'Q slot 为负' + assert self.q_pending_sem._value >= 0, 'Q pending 为负' + assert self.q_pending_sem._value <= self.q_pending_cap, 'Q pending 超限' + assert self.phys_sem._value >= 0, 'Physical 为负' + # inventory 深度不超过 slot 总量 + assert self.inventory.t_depth <= self.t_slot_cap, 'T inventory 超限' + assert self.inventory.q_depth <= self.q_slot_cap, 'Q inventory 超限' + + def check_invariants_final(self): + """最终严格守恒检查: 所有 Worker 停止后执行。 + + 此时没有 Worker 持有中间状态的 slot,也没有 PairLease。 + free + inventory == total + """ + self.conservation.check_t_conservation(self.t_slot_cap) + self.conservation.check_q_conservation(self.q_slot_cap) + self.conservation.check_pending_conservation(self.q_pending_cap) + + async def s_produce(self): + """S 产生一个 T。返回 True/False/None(cancelled)。""" + # 模拟物理资源使用 + await self.phys_sem.acquire() + delay = self.rng.uniform(0.001, 0.02) + try: + await asyncio.sleep(delay) + finally: + self.phys_sem.release() + + # 模拟随机失败 + if self.rng.random() < 0.1: + self.metrics.t_discarded += 1 + await self.events.record('t_discarded', 'gen_fail') + return False + + token = f'tok_{self.metrics.t_produced}' + self.metrics.t_produced += 1 + await self.events.record('t_produced', token) + + # 创建 envelope + 入库 + # 关键: create_with_slot 成功后 slot 在 envelope 中, + # 必须追踪 worker_held 直到 put_t 成功或 discard。 + now = time.time() + expire = now + self.t_max_age if self.rng.random() > 0.15 else now - 1 + t_env = None + try: + t_env = await ResourceEnvelope.create_with_slot( + 'T', token, self.t_slot_sem, expires_at=expire + ) + # slot 已获取,标记 worker 持有 + self.conservation.worker_acquire_t() + await self.inventory.put_t(t_env) + # put 成功,所有权转移到 inventory + self.conservation.worker_release_t() + except Exception: + if t_env is not None and not t_env.released: + t_env.discard(reason='error_cleanup') + self.conservation.worker_release_t() + self.metrics.t_discarded += 1 + return False + except asyncio.CancelledError: + if t_env is not None and not t_env.released: + t_env.discard(reason='cancel_cleanup') + self.conservation.worker_release_t() + self.metrics.t_discarded += 1 + raise + + self.check_invariants_soft() + return True + + async def p_produce(self): + """P 产生一个 Q。返回 True/False/None(cancelled)。""" + await self.q_pending_sem.acquire() + _pend_released = False + + try: + # 创建邮箱 + 发请求(持有 Physical) + await self.phys_sem.acquire() + try: + delay = self.rng.uniform(0.001, 0.02) + await asyncio.sleep(delay) + finally: + self.phys_sem.release() + + # 模拟随机失败 + if self.rng.random() < 0.1: + self.metrics.q_discarded += 1 + self.q_pending_sem.release() + _pend_released = True + await self.events.record('q_discarded', 'send_fail') + return False + + self.metrics.q_sent += 1 + + # 等待 Q 返回(不持有 Physical) + poll_delay = self.rng.uniform(0.001, 0.05) + await asyncio.sleep(poll_delay) + + # 模拟超时/失败 + if self.rng.random() < 0.1: + self.metrics.q_discarded += 1 + self.q_pending_sem.release() + _pend_released = True + await self.events.record('q_discarded', 'poll_timeout') + return False + + code = f'{self.rng.randint(100000, 999999)}' + self.metrics.q_returned += 1 + await self.events.record('q_returned', code) + + # 创建 envelope + 入库 + now = time.time() + expire = now + self.q_max_age if self.rng.random() > 0.15 else now - 1 + q_env = None + try: + q_env = await ResourceEnvelope.create_with_slot( + 'Q', {'email': f'e@{code}', 'password': 'p', 'code': code}, + self.q_slot_sem, expires_at=expire, + ) + self.conservation.worker_acquire_q() + await self.inventory.put_q(q_env) + self.conservation.worker_release_q() + except Exception: + if q_env is not None and not q_env.released: + q_env.discard(reason='error_cleanup') + self.conservation.worker_release_q() + self.metrics.q_discarded += 1 + except asyncio.CancelledError: + if q_env is not None and not q_env.released: + q_env.discard(reason='cancel_cleanup') + self.conservation.worker_release_q() + self.metrics.q_discarded += 1 + raise + + except asyncio.CancelledError: + if not _pend_released: + self.q_pending_sem.release() + _pend_released = True + raise + except Exception: + self.metrics.q_discarded += 1 + finally: + if not _pend_released: + self.q_pending_sem.release() + + self.check_invariants_soft() + return True + + async def c_consume(self): + """C 消费一个 pair。返回 True/False/None(cancelled)。""" + try: + async with self.inventory.claim_pair() as pair: + t_val = pair.t.value + q_val = pair.q.value + + # 检测重复 claim + assert t_val not in self._claimed_tokens, f'重复 claim: {t_val}' + self._claimed_tokens.add(t_val) + + # pairleased 计数 — 必须在 claim 后立即标记, + # 否则在下一个 await 前被取消会导致守恒违反 + self._pairleased_t += 1 + self._pairleased_q += 1 + + # 消费(持有 Physical) + await self.phys_sem.acquire() + try: + delay = self.rng.uniform(0.001, 0.02) + await asyncio.sleep(delay) + + # 模拟随机失败 + if self.rng.random() < 0.15: + self.metrics.pair_consumed_fail += 1 + await self.events.record('pair_fail', t_val) + else: + self.metrics.pair_consumed_ok += 1 + self.metrics.success_count += 1 + await self.events.record('pair_ok', t_val) + finally: + self.phys_sem.release() + + # PairLease 退出后(正常路径) + self._pairleased_t -= 1 + self._pairleased_q -= 1 + + except asyncio.CancelledError: + # PairLease.__aexit__ 已核销 slot。 + # 如果 _pairleased 已 +1,则 -1;否则不操作。 + if self._pairleased_t > 0: + self._pairleased_t -= 1 + if self._pairleased_q > 0: + self._pairleased_q -= 1 + raise + except Exception: + self.metrics.pair_consumed_fail += 1 + if self._pairleased_t > 0: + self._pairleased_t -= 1 + if self._pairleased_q > 0: + self._pairleased_q -= 1 + + self.check_invariants_soft() + return True + + +@pytest.mark.slow +class TestProperty: + """性质测试: 随机事件流 + 持续不变量检查。""" + + @pytest.mark.asyncio + async def test_random_events_10s(self): + """随机 S/P/C 事件流 10 秒,持续检查不变量。""" + engine = PropertyEngine(seed=42) + stop = asyncio.Event() + errors = [] + + async def s_loop(): + while not stop.is_set(): + try: + await engine.s_produce() + except asyncio.CancelledError: + break + except AssertionError as e: + errors.append(f'S: {e}') + break + await asyncio.sleep(engine.rng.uniform(0.005, 0.03)) + + async def p_loop(): + while not stop.is_set(): + try: + await engine.p_produce() + except asyncio.CancelledError: + break + except AssertionError as e: + errors.append(f'P: {e}') + break + await asyncio.sleep(engine.rng.uniform(0.005, 0.03)) + + async def c_loop(): + while not stop.is_set(): + try: + await engine.c_consume() + except asyncio.CancelledError: + break + except AssertionError as e: + errors.append(f'C: {e}') + break + await asyncio.sleep(engine.rng.uniform(0.005, 0.03)) + + tasks = [] + for _ in range(2): + tasks.append(asyncio.create_task(s_loop())) + tasks.append(asyncio.create_task(p_loop())) + tasks.append(asyncio.create_task(c_loop())) + + await asyncio.sleep(10) + stop.set() + for t in tasks: + t.cancel() + for t in tasks: + try: + await t + except asyncio.CancelledError: + pass + + # 最终不变量检查 + engine.check_invariants_final() + + assert not errors, f'不变量违反:\n' + '\n'.join(errors) + print(f'\n[Property] t_produced={engine.metrics.t_produced} ' + f'q_sent={engine.metrics.q_sent} q_returned={engine.metrics.q_returned} ' + f'pair_ok={engine.metrics.pair_consumed_ok} pair_fail={engine.metrics.pair_consumed_fail} ' + f'success={engine.metrics.success_count}') + + @pytest.mark.asyncio + async def test_random_with_cancellations(self): + """随机事件 + 随机取消 Worker,不变量仍成立。""" + engine = PropertyEngine(seed=123) + stop = asyncio.Event() + errors = [] + + async def s_loop(): + while not stop.is_set(): + try: + await engine.s_produce() + except asyncio.CancelledError: + break + except AssertionError as e: + errors.append(f'S: {e}') + break + await asyncio.sleep(engine.rng.uniform(0.005, 0.03)) + + async def p_loop(): + while not stop.is_set(): + try: + await engine.p_produce() + except asyncio.CancelledError: + break + except AssertionError as e: + errors.append(f'P: {e}') + break + await asyncio.sleep(engine.rng.uniform(0.005, 0.03)) + + async def c_loop(): + while not stop.is_set(): + try: + await engine.c_consume() + except asyncio.CancelledError: + break + except AssertionError as e: + errors.append(f'C: {e}') + break + await asyncio.sleep(engine.rng.uniform(0.005, 0.03)) + + async def cancel_loop(): + """随机取消并重启 Worker。""" + while not stop.is_set(): + await asyncio.sleep(engine.rng.uniform(0.5, 2.0)) + if tasks: + idx = engine.rng.randint(0, len(tasks) - 1) + t = tasks[idx] + if not t.done(): + t.cancel() + try: + await t + except asyncio.CancelledError: + pass + # 重启 + role = idx % 3 + if role == 0: + tasks[idx] = asyncio.create_task(s_loop()) + elif role == 1: + tasks[idx] = asyncio.create_task(p_loop()) + else: + tasks[idx] = asyncio.create_task(c_loop()) + + tasks = [] + for _ in range(2): + tasks.append(asyncio.create_task(s_loop())) + tasks.append(asyncio.create_task(p_loop())) + tasks.append(asyncio.create_task(c_loop())) + tasks.append(asyncio.create_task(cancel_loop())) + + await asyncio.sleep(10) + stop.set() + for t in tasks: + t.cancel() + for t in tasks: + try: + await t + except asyncio.CancelledError: + pass + + engine.check_invariants_final() + assert not errors, f'不变量违反:\n' + '\n'.join(errors) + print(f'\n[Property+Cancel] t_produced={engine.metrics.t_produced} ' + f'pair_ok={engine.metrics.pair_consumed_ok} ' + f'success={engine.metrics.success_count}') + + @pytest.mark.asyncio + async def test_extreme_ttl_pressure(self): + """极端 TTL: 所有资源几乎立即过期,验证过期清理不破坏守恒。""" + engine = PropertyEngine( + t_slot_cap=4, q_slot_cap=4, q_pending_cap=6, phys_cap=3, + t_max_age=0.05, q_max_age=0.05, # 50ms 过期 + seed=789, + ) + stop = asyncio.Event() + errors = [] + + async def s_loop(): + while not stop.is_set(): + try: + await engine.s_produce() + except asyncio.CancelledError: + break + except AssertionError as e: + errors.append(f'S: {e}') + break + await asyncio.sleep(engine.rng.uniform(0.01, 0.05)) + + async def p_loop(): + while not stop.is_set(): + try: + await engine.p_produce() + except asyncio.CancelledError: + break + except AssertionError as e: + errors.append(f'P: {e}') + break + await asyncio.sleep(engine.rng.uniform(0.01, 0.05)) + + async def c_loop(): + while not stop.is_set(): + try: + await engine.c_consume() + except asyncio.CancelledError: + break + except AssertionError as e: + errors.append(f'C: {e}') + break + await asyncio.sleep(engine.rng.uniform(0.01, 0.05)) + + tasks = [] + for _ in range(2): + tasks.append(asyncio.create_task(s_loop())) + tasks.append(asyncio.create_task(p_loop())) + tasks.append(asyncio.create_task(c_loop())) + + await asyncio.sleep(5) + stop.set() + for t in tasks: + t.cancel() + for t in tasks: + try: + await t + except asyncio.CancelledError: + pass + + engine.check_invariants_final() + assert not errors, f'不变量违反:\n' + '\n'.join(errors) + print(f'\n[TTL Pressure] t_produced={engine.metrics.t_produced} ' + f't_expired={engine.metrics.t_expired} q_expired={engine.metrics.q_expired} ' + f'pair_ok={engine.metrics.pair_consumed_ok}') diff --git a/tests/test_register_runtime_unittest.py b/tests/test_register_runtime_unittest.py new file mode 100644 index 0000000..44b70ce --- /dev/null +++ b/tests/test_register_runtime_unittest.py @@ -0,0 +1,1325 @@ +import asyncio +import json +import sys +import tempfile +import types +import unittest + +from grok_register.core.observer import Metrics + + +playwright_pkg = types.ModuleType("playwright") +playwright_async_api = types.ModuleType("playwright.async_api") +playwright_async_api.async_playwright = lambda: None +sys.modules.setdefault("playwright", playwright_pkg) +sys.modules.setdefault("playwright.async_api", playwright_async_api) + +requests_mod = types.ModuleType("requests") +requests_mod.get = lambda *_args, **_kwargs: None +requests_mod.post = lambda *_args, **_kwargs: None +sys.modules.setdefault("requests", requests_mod) + +from grok_register import register + + +def test_registration_persists_exact_session_before_legacy_outputs(monkeypatch): + writes = [] + monkeypatch.setattr( + register, + "_append_registration_line", + lambda path, line, mode=None, durable=False: writes.append( + (path, line, mode, durable) + ), + ) + + register._persist_registration( + "user@example.test", + "password", + "sso-token", + [{"name": "sso", "value": "opaque", "domain": "accounts.x.ai"}], + ) + + assert [item[0] for item in writes] == [ + "keys/auth-sessions.jsonl", + "keys/accounts.txt", + "keys/grok.txt", + ] + assert writes[0][2] == 0o600 + assert [item[3] for item in writes] == [True, False, False] + + +class FakePage: + def __init__(self, context=None): + self.context = context + self.closed = False + self.url = "about:blank" + self.goto_calls = [] + self.waits = [] + self.evaluations = [] + self.route_calls = [] + self.mouse = types.SimpleNamespace( + moves=[], + clicks=[], + downs=0, + ups=0, + move=self._mouse_move, + click=self._mouse_click, + down=self._mouse_down, + up=self._mouse_up, + ) + self.turnstile_token = "" + self.turnstile_box = {"x": 160, "y": 45} + self.turnstile_page_trace = { + "created_at": 10.0, + "script_inserted_at": 11.0, + "script_loaded_at": 12.0, + "render_called_at": 13.0, + "render_returned_at": 14.0, + "token_written_at": None, + "token_len": 0, + "error": None, + } + self.turnstile_dom_snapshot = { + "widget": {"present": True, "x": 10, "y": 10, "w": 300, "h": 70, "visible": True}, + "click_center": {"x": 160, "y": 45}, + "element_at_center": {"tag": "IFRAME", "id": "", "class": "", "is_iframe": True}, + "all_iframe_count": 1, + "turnstile_iframe_count": 1, + "iframe_summaries": [ + {"host": "challenges.cloudflare.com", "path": "/turnstile/v0", "x": 10, "y": 10, "w": 300, "h": 70, "visible": True} + ], + "turnstile_loaded": True, + "response_input": {"present": True, "token_len": 0}, + } + + async def set_viewport_size(self, size): + self.viewport = size + pass + + async def goto(self, url, timeout=None, wait_until=None): + self.url = url + self.goto_calls.append({ + "url": url, + "timeout": timeout, + "wait_until": wait_until, + }) + pass + + async def wait_for_timeout(self, timeout): + self.waits.append(timeout) + pass + + async def evaluate(self, script): + self.evaluations.append(script) + if "__csp_solver_snapshot" in script: + return self.turnstile_dom_snapshot + if "__cspTurnstileTrace" in script and "return window.__cspTurnstileTrace" in script: + return self.turnstile_page_trace + if "cf-turnstile-response" in script: + return self.turnstile_token + if "getBoundingClientRect" in script and ".cf-turnstile" in script: + return self.turnstile_box + return None + + async def close(self): + self.closed = True + pass + + async def route(self, pattern, handler): + self.route_calls.append({"pattern": pattern, "handler": handler}) + + async def _mouse_move(self, x, y, steps=None): + self.mouse.moves.append({"x": x, "y": y, "steps": steps}) + + async def _mouse_click(self, x, y): + self.mouse.clicks.append({"x": x, "y": y}) + + async def _mouse_down(self): + self.mouse.downs += 1 + + async def _mouse_up(self): + self.mouse.ups += 1 + + +class FakeContext: + def __init__(self): + self.pages = [] + self.closed = False + self.clear_cookies_calls = 0 + self.request = types.SimpleNamespace(get=self._request_get) + self.request_get_calls = [] + self.cookies_value = [] + self.cancel_on_clear = False + + async def new_page(self): + page = FakePage(self) + self.pages.append(page) + return page + + async def clear_cookies(self): + self.clear_cookies_calls += 1 + if self.cancel_on_clear: + raise asyncio.CancelledError() + self.cookies_value = [] + pass + + async def cookies(self): + return list(self.cookies_value) + + async def _request_get(self, url, timeout=None): + self.request_get_calls.append({"url": url, "timeout": timeout}) + return types.SimpleNamespace(status=403) + + async def close(self): + self.closed = True + for page in self.pages: + page.closed = True + + +class FakeBrowser: + def __init__(self): + self.pages = [] + self.contexts = [] + self.context = types.SimpleNamespace(request=object()) + + async def new_page(self): + page = FakePage(self.context) + self.pages.append(page) + return page + + async def new_context(self): + context = FakeContext() + self.contexts.append(context) + return context + + +class FakePair: + def __init__(self): + self.t = types.SimpleNamespace(value="tok") + self.q = types.SimpleNamespace( + value={"email": "e@example.test", "password": "pw", "code": "123456"} + ) + + +class FakeInventory: + def __init__(self): + self.active = 0 + self.claims = 0 + self.t_depth = 0 + self.q_depth = 0 + + def claim_pair(self): + inventory = self + + class Claim: + async def __aenter__(self): + inventory.active += 1 + inventory.claims += 1 + return FakePair() + + async def __aexit__(self, exc_type, exc, tb): + inventory.active -= 1 + return False + + return Claim() + + +class RegisterRuntimeTests(unittest.IsolatedAsyncioTestCase): + async def asyncSetUp(self): + self._old_stop = register.STOP + self._old_timeout = getattr(register, "C_CONSUME_TIMEOUT", None) + self._old_verify = register.grpc_verify_code + self._old_register = register.server_action_register + self._old_create_code = register.grpc_create_code + self._old_poll_code = register.poll_code + self._old_poll_code_async = register._poll_code_async + self._old_log = register.log + self._old_target = register.TARGET + self._old_c_hot_page_pool = getattr(register, "C_HOT_PAGE_POOL", None) + self._old_c_hot_page_pool_size = getattr(register, "C_HOT_PAGE_POOL_SIZE", None) + self._old_c_set_cookie_via_request = getattr(register, "C_SET_COOKIE_VIA_REQUEST", None) + self._old_log_mode = getattr(register, "REGISTER_LOG_MODE", None) + self._old_rate_limit_circuit = register.REGISTRATION_RATE_LIMIT_CIRCUIT + + async def asyncTearDown(self): + if hasattr(register, "_close_c_hot_page_pool"): + await register._close_c_hot_page_pool() + register.STOP = self._old_stop + if self._old_timeout is not None: + register.C_CONSUME_TIMEOUT = self._old_timeout + register.grpc_verify_code = self._old_verify + register.server_action_register = self._old_register + register.grpc_create_code = self._old_create_code + register.poll_code = self._old_poll_code + register._poll_code_async = self._old_poll_code_async + register.log = self._old_log + register.TARGET = self._old_target + if self._old_c_hot_page_pool is None and hasattr(register, "C_HOT_PAGE_POOL"): + delattr(register, "C_HOT_PAGE_POOL") + elif self._old_c_hot_page_pool is not None: + register.C_HOT_PAGE_POOL = self._old_c_hot_page_pool + if self._old_c_hot_page_pool_size is None and hasattr(register, "C_HOT_PAGE_POOL_SIZE"): + delattr(register, "C_HOT_PAGE_POOL_SIZE") + elif self._old_c_hot_page_pool_size is not None: + register.C_HOT_PAGE_POOL_SIZE = self._old_c_hot_page_pool_size + if self._old_c_set_cookie_via_request is None and hasattr(register, "C_SET_COOKIE_VIA_REQUEST"): + delattr(register, "C_SET_COOKIE_VIA_REQUEST") + elif self._old_c_set_cookie_via_request is not None: + register.C_SET_COOKIE_VIA_REQUEST = self._old_c_set_cookie_via_request + if self._old_log_mode is None and hasattr(register, "REGISTER_LOG_MODE"): + delattr(register, "REGISTER_LOG_MODE") + elif self._old_log_mode is not None: + register.REGISTER_LOG_MODE = self._old_log_mode + register.REGISTRATION_RATE_LIMIT_CIRCUIT = self._old_rate_limit_circuit + + async def test_c_worker_timeout_releases_physical_and_pair_and_counts_failure(self): + async def slow_verify(*_args, **_kwargs): + await asyncio.sleep(10) + return True + + register.STOP = asyncio.Event() + register.C_CONSUME_TIMEOUT = 0.05 + register.grpc_verify_code = slow_verify + register.server_action_register = lambda *_args, **_kwargs: None + register.log = lambda _msg: None + + inventory = FakeInventory() + physical_sem = asyncio.Semaphore(1) + metrics = Metrics() + + task = asyncio.create_task( + register.c_worker(0, FakeBrowser(), inventory, physical_sem, metrics) + ) + await asyncio.sleep(0.12) + register.STOP.set() + await asyncio.wait_for(task, timeout=1) + + self.assertEqual(physical_sem._value, 1) + self.assertEqual(inventory.active, 0) + self.assertEqual(metrics.pair_consumed_fail, 1) + + async def test_c_consume_uses_single_use_page_by_default(self): + async def ok_verify(*_args, **_kwargs): + return True + + async def no_sso_register(*_args, **_kwargs): + return None + + register.C_HOT_PAGE_POOL = False + register.grpc_verify_code = ok_verify + register.server_action_register = no_sso_register + register.log = lambda _msg: None + + browser = FakeBrowser() + physical_sem = asyncio.Semaphore(1) + + ok = await register._consume_pair(browser, physical_sem, FakePair(), Metrics()) + + self.assertFalse(ok) + self.assertEqual(len(browser.pages), 1) + self.assertEqual(len(browser.contexts), 0) + self.assertTrue(browser.pages[0].closed) + self.assertEqual(physical_sem._value, 1) + + async def test_c_hot_page_reuses_page_and_clears_cookies_between_consumes(self): + seen_pages = [] + + async def ok_verify(*_args, **_kwargs): + return True + + async def no_sso_register(page, *_args, **_kwargs): + seen_pages.append(page) + return None + + register.C_HOT_PAGE_POOL = True + register.C_HOT_PAGE_POOL_SIZE = 2 + register.grpc_verify_code = ok_verify + register.server_action_register = no_sso_register + register.log = lambda _msg: None + + browser = FakeBrowser() + physical_sem = asyncio.Semaphore(1) + + first = await register._consume_pair(browser, physical_sem, FakePair(), Metrics()) + second = await register._consume_pair(browser, physical_sem, FakePair(), Metrics()) + + self.assertFalse(first) + self.assertFalse(second) + self.assertEqual(len(browser.contexts), 1) + self.assertEqual(len(browser.contexts[0].pages), 1) + self.assertEqual(seen_pages, [browser.contexts[0].pages[0], browser.contexts[0].pages[0]]) + self.assertEqual(browser.contexts[0].clear_cookies_calls, 2) + self.assertFalse(browser.contexts[0].closed) + self.assertFalse(browser.contexts[0].pages[0].closed) + self.assertEqual(physical_sem._value, 1) + + async def test_c_hot_page_discards_page_after_exception(self): + async def failing_verify(*_args, **_kwargs): + raise RuntimeError("verify failed") + + register.C_HOT_PAGE_POOL = True + register.C_HOT_PAGE_POOL_SIZE = 2 + register.grpc_verify_code = failing_verify + register.log = lambda _msg: None + + browser = FakeBrowser() + physical_sem = asyncio.Semaphore(1) + + with self.assertRaises(RuntimeError): + await register._consume_pair(browser, physical_sem, FakePair(), Metrics()) + + self.assertEqual(len(browser.contexts), 1) + self.assertTrue(browser.contexts[0].closed) + self.assertTrue(browser.contexts[0].pages[0].closed) + self.assertEqual(physical_sem._value, 1) + + async def test_c_hot_page_closes_context_when_cancelled_during_cleanup(self): + context = FakeContext() + page = await context.new_page() + page.url = "https://accounts.x.ai/sign-up?redirect=grok-com" + context.cancel_on_clear = True + register.C_HOT_PAGE_POOL = True + + with self.assertRaises(asyncio.CancelledError): + await register._release_c_page(context, page, healthy=True) + + self.assertTrue(context.closed) + self.assertTrue(page.closed) + + async def test_c_hot_page_discards_page_after_fallback_navigation(self): + context = FakeContext() + page = await context.new_page() + page.url = "https://example.test/set-cookie?q=abc" + register.C_HOT_PAGE_POOL = True + register.C_HOT_PAGE_POOL_SIZE = 2 + + await register._release_c_page(context, page, healthy=True) + + self.assertTrue(context.closed) + self.assertTrue(page.closed) + self.assertEqual(register._c_hot_page_pool, []) + + async def test_server_action_can_set_cookie_via_request_without_navigating_page(self): + page = FakePage() + context = FakeContext() + context.cookies_value = [{"name": "sso", "value": "x" * 152}] + page.context = context + page.goto_calls = [] + + async def fake_evaluate(_script): + return '0:"https:\\/\\/auth.grokipedia.com\\/set-cookie?q=abc"1:' + + page.evaluate = fake_evaluate + register.STATE_TREE = "state" + register.ACTION_ID = "action" + register.C_SET_COOKIE_VIA_REQUEST = True + + sso = await register.server_action_register( + page, "e@example.test", "pw", "123456", "token" + ) + + self.assertEqual(sso, "x" * 152) + self.assertEqual(page.goto_calls, []) + self.assertEqual( + context.request_get_calls, + [{"url": "https://auth.grokipedia.com/set-cookie?q=abc", "timeout": 15000}], + ) + + async def test_server_action_raises_a_distinct_error_for_a_rate_limited_signup_page(self): + page = FakePage() + + async def rate_limited_response(_script): + return "Too many requests. Please try again later." + + page.evaluate = rate_limited_response + register.STATE_TREE = "state" + register.ACTION_ID = "action" + + with self.assertRaises(register.RegistrationRateLimited): + await register.server_action_register( + page, "e@example.test", "pw", "123456", "token" + ) + + async def test_rate_limit_circuit_opens_for_the_configured_cooldown(self): + now = [100.0] + circuit = register.RegistrationRateLimitCircuit( + cooldown_seconds=60, + clock=lambda: now[0], + ) + + circuit.trip() + + self.assertTrue(circuit.is_open()) + self.assertEqual(circuit.remaining_seconds(), 60) + now[0] = 160.0 + self.assertFalse(circuit.is_open()) + + async def test_monitor_uses_metrics_snapshot(self): + register.STOP = asyncio.Event() + register.TARGET = 1 + register.REGISTER_LOG_MODE = "debug" + messages = [] + register.log = messages.append + + metrics = Metrics() + metrics.success_count = 1 + metrics.pair_claimed = 2 + metrics.pair_consumed_ok = 1 + metrics.pair_consumed_fail = 1 + sems = { + "physical": asyncio.Semaphore(1), + "t_slot": asyncio.Semaphore(1), + "q_slot": asyncio.Semaphore(1), + "q_pending": asyncio.Semaphore(1), + } + + await register.monitor(FakeInventory(), sems, metrics, interval=0) + + self.assertTrue( + any("pair:2 ok:1 fail:1" in message for message in messages), + messages, + ) + + async def test_monitor_hides_internal_snapshot_in_user_mode(self): + register.STOP = asyncio.Event() + register.TARGET = 1 + register.REGISTER_LOG_MODE = "user" + messages = [] + register.log = messages.append + + metrics = Metrics() + metrics.success_count = 1 + sems = { + "physical": asyncio.Semaphore(1), + "t_slot": asyncio.Semaphore(1), + "q_slot": asyncio.Semaphore(1), + "q_pending": asyncio.Semaphore(1), + } + + await register.monitor(FakeInventory(), sems, metrics, interval=0) + + self.assertFalse(any(message.startswith("[*] T:") for message in messages), messages) + + async def test_user_event_format_reports_only_registration_outcomes(self): + self.assertEqual( + register.format_user_registration_event("started", task_id=7), + "[→] 开始注册 #7", + ) + self.assertEqual( + register.format_user_registration_event( + "success", task_id=7, count=5, rate_per_minute=12.34 + ), + "[✓] 注册成功 #7 | 运行平均 12.3/分 | 累计 5", + ) + self.assertEqual( + register.format_user_registration_event("failed", task_id=7), + "[✗] 注册失败 #7 | 已跳过,继续下一任务", + ) + self.assertEqual( + register.format_user_registration_event("rate_limited", wait_seconds=60), + "[⏸] 触发限流 | 60秒后恢复探测", + ) + self.assertEqual( + register.format_user_registration_event("recovered", wait_seconds=61), + "[▶] 限流解除 | 实际等待 61秒", + ) + + async def test_debug_flag_overrides_environment_and_invalid_mode_is_rejected(self): + self.assertEqual( + register.resolve_register_log_mode(["--debug"], {"REGISTER_LOG_MODE": "user"}), + "debug", + ) + self.assertEqual( + register.resolve_register_log_mode([], {"REGISTER_LOG_MODE": "debug"}), + "debug", + ) + with self.assertRaises(ValueError): + register.resolve_register_log_mode([], {"REGISTER_LOG_MODE": "verbose"}) + + async def test_registration_task_numbers_are_separate_from_pair_claims(self): + metrics = Metrics() + metrics.pair_claimed = 12 + + self.assertEqual(metrics.next_registration_task(), 1) + self.assertEqual(metrics.next_registration_task(), 2) + + async def test_five_minute_rate_uses_process_uptime_then_sliding_window(self): + now = [0.0] + metrics = Metrics(clock=lambda: now[0]) + self.assertIsNone(metrics.five_minute_success_rate()) + + now[0] = 10.0 + metrics.record_success() + now[0] = 20.0 + metrics.record_success() + self.assertEqual(metrics.five_minute_success_rate(), 6.0) + + now[0] = 311.0 + self.assertAlmostEqual(metrics.five_minute_success_rate(), 0.2) + now[0] = 321.0 + self.assertEqual(metrics.five_minute_success_rate(), 0.0) + + async def test_runtime_average_rate_includes_the_entire_cooldown_period(self): + now = [0.0] + metrics = Metrics(clock=lambda: now[0]) + self.assertIsNone(metrics.runtime_average_success_rate()) + + now[0] = 10.0 + metrics.record_success() + now[0] = 20.0 + metrics.record_success() + self.assertEqual(metrics.runtime_average_success_rate(), 6.0) + + now[0] = 320.0 + self.assertEqual(metrics.runtime_average_success_rate(), 0.375) + + async def test_terminal_output_failure_does_not_escape_log(self): + old_output = register._terminal_output + try: + register._terminal_output = lambda _message: (_ for _ in ()).throw( + OSError("closed pipe") + ) + register.log("safe") + finally: + register._terminal_output = old_output + + async def test_rate_limit_circuit_measures_one_recovery_window(self): + now = [100.0] + circuit = register.RegistrationRateLimitCircuit( + cooldown_seconds=60, + clock=lambda: now[0], + ) + circuit.trip() + now[0] = 161.5 + + self.assertEqual(circuit.consume_recovery_seconds(), 61.5) + self.assertIsNone(circuit.consume_recovery_seconds()) + + async def test_rate_limit_probe_is_released_when_consume_times_out(self): + async def slow_verify(*_args, **_kwargs): + await asyncio.Event().wait() + + register.REGISTRATION_RATE_LIMIT_CIRCUIT = register.RegistrationRateLimitCircuit(0) + register.REGISTRATION_RATE_LIMIT_CIRCUIT.trip() + register.grpc_verify_code = slow_verify + + with self.assertRaises(asyncio.TimeoutError): + await asyncio.wait_for( + register._consume_pair( + FakeBrowser(), asyncio.Semaphore(1), FakePair(), Metrics(), task_id=1 + ), + timeout=0.01, + ) + + self.assertFalse(register.REGISTRATION_RATE_LIMIT_CIRCUIT._probe_active) + + async def test_auto_capacity_is_bounded_by_cpu_and_memory(self): + roomy = register.derive_capacity( + cpu_count=2, + max_mem_mb=5600, + physical_cap=0, + physical_per_cpu=4, + physical_mem_mb=512, + min_free_mem_mb=500, + ) + tight = register.derive_capacity( + cpu_count=2, + max_mem_mb=1100, + physical_cap=0, + physical_per_cpu=4, + physical_mem_mb=512, + min_free_mem_mb=500, + ) + + self.assertEqual(roomy[0], 8) + self.assertEqual(roomy[1], 10) + self.assertEqual(roomy[2], register.Q_PENDING_CAP + 2) + self.assertEqual(roomy[3], 10) + self.assertEqual(tight[0], 1) + + async def test_default_auto_capacity_is_conservative(self): + physical, s_workers, _p_workers, c_workers = register.derive_capacity( + cpu_count=2, + max_mem_mb=5600, + physical_cap=0, + ) + + self.assertEqual(physical, 4) + self.assertEqual(s_workers, 6) + self.assertEqual(c_workers, 6) + + async def test_explicit_physical_cap_overrides_auto_capacity(self): + physical, s_workers, _p_workers, c_workers = register.derive_capacity( + cpu_count=2, + max_mem_mb=1100, + physical_cap=3, + physical_per_cpu=4, + physical_mem_mb=512, + min_free_mem_mb=500, + ) + + self.assertEqual(physical, 3) + self.assertEqual(s_workers, 5) + self.assertEqual(c_workers, 5) + + async def test_capacity_profile_supplies_physical_cap_when_not_explicit(self): + physical, s_workers, _p_workers, c_workers = register.derive_capacity( + cpu_count=2, + max_mem_mb=5600, + physical_cap=0, + profile_physical_cap=7, + ) + tight, *_ = register.derive_capacity( + cpu_count=2, + max_mem_mb=1100, + physical_cap=0, + profile_physical_cap=7, + ) + explicit, *_ = register.derive_capacity( + cpu_count=2, + max_mem_mb=5600, + physical_cap=5, + profile_physical_cap=7, + ) + + self.assertEqual(physical, 7) + self.assertEqual(s_workers, 9) + self.assertEqual(c_workers, 9) + self.assertEqual(tight, 1) + self.assertEqual(explicit, 5) + + async def test_load_capacity_profile_reads_valid_physical_cap(self): + with tempfile.NamedTemporaryFile("w+", delete=True) as f: + json.dump({"physical_cap": 7}, f) + f.flush() + + profile = register.load_capacity_profile(f.name) + + self.assertEqual(profile["physical_cap"], 7) + self.assertEqual(register.load_capacity_profile("/does/not/exist"), {}) + self.assertEqual(register.load_capacity_profile(""), {}) + + async def test_admission_t_high_defaults_to_physical_cap_bounded_by_slot(self): + watermarks = register.derive_admission_watermarks( + physical_cap=6, + t_slot_cap=8, + q_pending_cap=12, + t_target=4, + q_target=4, + ) + bounded = register.derive_admission_watermarks( + physical_cap=10, + t_slot_cap=8, + q_pending_cap=12, + t_target=4, + q_target=4, + ) + + self.assertEqual(watermarks["t_high"], 6) + self.assertEqual(watermarks["t_low"], 3) + self.assertEqual(bounded["t_high"], 8) + + async def test_admission_t_high_override_remains_explicit(self): + watermarks = register.derive_admission_watermarks( + physical_cap=6, + t_slot_cap=8, + q_pending_cap=12, + t_target=4, + q_target=4, + t_high_override=4, + t_low_override=2, + ) + + self.assertEqual(watermarks["t_high"], 4) + self.assertEqual(watermarks["t_low"], 2) + + async def test_c_hot_page_pool_size_is_derived_at_startup(self): + self.assertEqual( + register.derive_c_hot_page_pool_size( + physical_cap=6, c_workers=8, configured_size=0 + ), + 6, + ) + self.assertEqual( + register.derive_c_hot_page_pool_size( + physical_cap=8, c_workers=3, configured_size=0 + ), + 3, + ) + self.assertEqual( + register.derive_c_hot_page_pool_size( + physical_cap=8, c_workers=10, configured_size=4 + ), + 4, + ) + + async def test_send_q_request_batch_reuses_one_page_for_multiple_emails(self): + emails = [] + + async def fake_create_code(_page, email): + emails.append(email) + return True + + register.grpc_create_code = fake_create_code + browser = FakeBrowser() + physical_sem = asyncio.Semaphore(1) + p_send_sem = asyncio.Semaphore(1) + requests = [ + {"handle": "h1", "email": "a@example.test", "password": "pw1"}, + {"handle": "h2", "email": "b@example.test", "password": "pw2"}, + {"handle": "h3", "email": "c@example.test", "password": "pw3"}, + ] + + results = await register._send_q_request_batch( + browser, physical_sem, p_send_sem, requests + ) + + self.assertEqual(emails, [item["email"] for item in requests]) + self.assertEqual([item["sent"] for item in results], [True, True, True]) + self.assertEqual(len(browser.pages), 1) + self.assertTrue(browser.pages[0].closed) + self.assertEqual(physical_sem._value, 1) + self.assertEqual(p_send_sem._value, 1) + + async def test_send_q_request_batch_records_physical_and_stage_metrics(self): + async def fake_create_code(_page, _email): + return True + + register.grpc_create_code = fake_create_code + browser = FakeBrowser() + physical_sem = asyncio.Semaphore(1) + p_send_sem = asyncio.Semaphore(1) + metrics = Metrics() + + await register._send_q_request_batch( + browser, + physical_sem, + p_send_sem, + [{"handle": "h1", "email": "a@example.test", "password": "pw"}], + metrics, + ) + + self.assertEqual(metrics.p_physical_count, 1) + self.assertEqual(metrics.p_page_prepare_count, 1) + self.assertEqual(metrics.p_send_count, 1) + self.assertGreaterEqual(metrics.p_physical_wait_seconds, 0) + self.assertGreaterEqual(metrics.p_physical_hold_seconds, 0) + + async def test_consume_pair_records_physical_and_stage_metrics(self): + async def ok_verify(*_args, **_kwargs): + return True + + async def no_sso_register(*_args, **_kwargs): + return None + + register.C_HOT_PAGE_POOL = False + register.grpc_verify_code = ok_verify + register.server_action_register = no_sso_register + register.log = lambda _msg: None + metrics = Metrics() + + await register._consume_pair( + FakeBrowser(), asyncio.Semaphore(1), FakePair(), metrics + ) + + self.assertEqual(metrics.c_physical_count, 1) + self.assertEqual(metrics.c_page_acquire_count, 1) + self.assertEqual(metrics.c_verify_count, 1) + self.assertEqual(metrics.c_register_count, 1) + self.assertEqual(metrics.c_hot_page_hits, 0) + self.assertEqual(metrics.c_hot_page_misses, 0) + + async def test_poll_and_admit_q_releases_one_pending_per_terminal_request(self): + register.P_REQUEST_TIMEOUT = 1 + register.poll_code = lambda _handle: None + q_pending_sem = asyncio.Semaphore(0) + q_slot_sem = asyncio.Semaphore(1) + metrics = Metrics() + + await register._poll_and_admit_q( + {"handle": "h1", "email": "a@example.test", "password": "pw"}, + FakeInventory(), + q_pending_sem, + q_slot_sem, + metrics, + ) + + self.assertEqual(q_pending_sem._value, 1) + self.assertEqual(q_slot_sem._value, 1) + self.assertEqual(metrics.q_discarded, 1) + + async def test_poll_cancel_before_terminal_does_not_release_pending(self): + async def blocked_poll(_loop, _handle): + await asyncio.sleep(10) + return "123456" + + register._poll_code_async = blocked_poll + q_pending_sem = asyncio.Semaphore(0) + q_slot_sem = asyncio.Semaphore(1) + metrics = Metrics() + + task = asyncio.create_task( + register._poll_and_admit_q( + {"handle": "h1", "email": "a@example.test", "password": "pw"}, + FakeInventory(), + q_pending_sem, + q_slot_sem, + metrics, + ) + ) + await asyncio.sleep(0.02) + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + + self.assertEqual(q_pending_sem._value, 0) + self.assertEqual(q_slot_sem._value, 1) + + async def test_poll_cancel_after_q_return_releases_pending(self): + async def returned_poll(_loop, _handle): + return "123456" + + class BlockingInventory(FakeInventory): + async def put_q(self, _env): + await asyncio.sleep(10) + + register._poll_code_async = returned_poll + q_pending_sem = asyncio.Semaphore(0) + q_slot_sem = asyncio.Semaphore(1) + metrics = Metrics() + + task = asyncio.create_task( + register._poll_and_admit_q( + {"handle": "h1", "email": "a@example.test", "password": "pw"}, + BlockingInventory(), + q_pending_sem, + q_slot_sem, + metrics, + ) + ) + await asyncio.sleep(0.02) + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + + self.assertEqual(q_pending_sem._value, 1) + self.assertEqual(q_slot_sem._value, 1) + self.assertEqual(metrics.q_returned, 1) + + async def test_metrics_snapshot_includes_solver_timing(self): + metrics = Metrics() + metrics.t_solve_count = 2 + metrics.t_solve_seconds = 5.0 + metrics.t_solve_failed = 1 + sems = { + "physical": asyncio.Semaphore(1), + "t_slot": asyncio.Semaphore(1), + "q_slot": asyncio.Semaphore(1), + "q_pending": asyncio.Semaphore(1), + } + + row = metrics.snapshot(FakeInventory(), sems) + + self.assertIn("t_solve_avg:2.5", row) + self.assertIn("t_solve_fail:1", row) + + async def test_metrics_snapshot_includes_role_physical_and_stage_timing(self): + metrics = Metrics() + metrics.s_physical_count = 2 + metrics.s_physical_wait_seconds = 1.0 + metrics.s_physical_hold_seconds = 6.0 + metrics.p_physical_count = 1 + metrics.p_physical_wait_seconds = 0.2 + metrics.p_physical_hold_seconds = 1.4 + metrics.c_physical_count = 3 + metrics.c_physical_wait_seconds = 0.9 + metrics.c_physical_hold_seconds = 7.5 + metrics.p_email_create_count = 2 + metrics.p_email_create_seconds = 1.0 + metrics.p_page_prepare_count = 1 + metrics.p_page_prepare_seconds = 0.8 + metrics.p_send_count = 1 + metrics.p_send_seconds = 0.4 + metrics.c_page_acquire_count = 2 + metrics.c_page_acquire_seconds = 0.6 + metrics.c_verify_count = 2 + metrics.c_verify_seconds = 0.8 + metrics.c_register_count = 2 + metrics.c_register_seconds = 3.0 + metrics.c_hot_page_hits = 4 + metrics.c_hot_page_misses = 1 + sems = { + "physical": asyncio.Semaphore(1), + "t_slot": asyncio.Semaphore(1), + "q_slot": asyncio.Semaphore(1), + "q_pending": asyncio.Semaphore(1), + } + + row = metrics.snapshot(FakeInventory(), sems) + + self.assertIn("s_phys:0.50/3.00", row) + self.assertIn("p_phys:0.20/1.40", row) + self.assertIn("c_phys:0.30/2.50", row) + self.assertIn("p_stage:0.50/0.80/0.40", row) + self.assertIn("c_stage:0.30/0.40/1.50", row) + self.assertIn("c_hot:4/1", row) + + async def test_record_solver_trace_accumulates_stage_metrics(self): + metrics = Metrics() + + register._record_solver_trace( + metrics, + { + "goto_s": 1.0, + "inject_s": 0.2, + "initial_s": 0.5, + "click_s": 0.1, + "wait_s": 20.0, + "reused": True, + "visible_frame": False, + }, + 21.8, + "token", + ) + register._record_solver_trace(metrics, {}, 10.0, None) + + self.assertEqual(metrics.t_solve_count, 2) + self.assertEqual(metrics.t_solve_failed, 1) + self.assertAlmostEqual(metrics.solver_goto_seconds, 1.0) + self.assertAlmostEqual(metrics.solver_wait_seconds, 20.0) + self.assertEqual(metrics.solver_reused_count, 1) + self.assertEqual(metrics.solver_visible_frame_count, 0) + + async def test_solver_timeline_records_click_and_token_events_when_enabled(self): + page = FakePage() + original_click = register._mouse_click_turnstile_center_trace + original_read = register._read_turnstile_token + clicked_once = False + timeline = register._new_solver_timeline(enabled=True) + + async def fake_click(_page, **_kwargs): + nonlocal clicked_once + clicked_once = True + return True, {"box_eval_ms": 0.0} + + async def fake_read(_page): + return "token-value-long" if clicked_once else "" + + try: + register._mouse_click_turnstile_center_trace = fake_click + register._read_turnstile_token = fake_read + + clicked = await register._repeat_mouse_click_turnstile(page, timeline=timeline) + finally: + register._mouse_click_turnstile_center_trace = original_click + register._read_turnstile_token = original_read + + events = [event["event"] for event in timeline["events"]] + self.assertTrue(clicked) + self.assertIn("click_before", events) + self.assertIn("click_after", events) + self.assertTrue(any(event.get("dom", {}).get("widget", {}).get("present") for event in timeline["events"])) + self.assertTrue(any(event.get("click_call_ms", 0) >= 0 for event in timeline["events"])) + self.assertTrue(any(event.get("token_len", 0) > 10 for event in timeline["events"])) + + async def test_turnstile_dom_snapshot_reports_click_target_without_text_or_full_urls(self): + page = FakePage() + + snapshot = await register._turnstile_dom_snapshot(page) + + self.assertEqual(snapshot["all_iframe_count"], 1) + self.assertEqual(snapshot["turnstile_iframe_count"], 1) + self.assertEqual(snapshot["widget"]["w"], 300) + self.assertEqual(snapshot["element_at_center"]["tag"], "IFRAME") + self.assertEqual(snapshot["iframe_summaries"][0]["host"], "challenges.cloudflare.com") + self.assertNotIn("src", snapshot["iframe_summaries"][0]) + self.assertNotIn("text", snapshot["element_at_center"]) + + async def test_mouse_click_turnstile_center_can_return_timing_trace(self): + page = FakePage() + + clicked, trace = await register._mouse_click_turnstile_center_trace(page) + + self.assertTrue(clicked) + self.assertEqual(trace["click_x"], 160.0) + self.assertEqual(trace["click_y"], 45.0) + for key in ("box_eval_ms", "mouse_move1_ms", "mouse_move2_ms", "mouse_down_ms", "mouse_up_ms"): + self.assertIn(key, trace) + self.assertGreaterEqual(trace[key], 0) + + async def test_inject_turnstile_widget_leaves_default_script_uninstrumented(self): + page = FakePage() + + await register._inject_turnstile_widget(page) + + script = page.evaluations[-1] + self.assertNotIn("__cspTurnstileTrace", script) + + async def test_inject_turnstile_widget_records_page_timeline_when_enabled(self): + page = FakePage() + + await register._inject_turnstile_widget(page, timeline=True) + + script = page.evaluations[-1] + self.assertIn("__cspTurnstileTrace", script) + self.assertIn("script_inserted_at", script) + self.assertIn("render_called_at", script) + self.assertIn("token_written_at", script) + + async def test_start_turnstile_challenge_records_page_trace_when_timeline_enabled(self): + browser = FakeBrowser() + messages = [] + old_trace = register.SOLVER_TIMELINE_TRACE + old_sample = register.SOLVER_TIMELINE_SAMPLE + old_emitted = register._solver_timeline_emitted + old_log = register.log + try: + register.SOLVER_TIMELINE_TRACE = True + register.SOLVER_TIMELINE_SAMPLE = 1 + register._solver_timeline_emitted = 0 + register.log = messages.append + + item = await register._start_turnstile_challenge(browser, fast_click=True) + await register._put_solver_page(item, False) + finally: + register.SOLVER_TIMELINE_TRACE = old_trace + register.SOLVER_TIMELINE_SAMPLE = old_sample + register._solver_timeline_emitted = old_emitted + register.log = old_log + + events = item["timeline"]["events"] + self.assertTrue(any(event["event"] == "page_trace_after_inject" for event in events)) + self.assertTrue(any(event["event"] == "page_trace_after_click" for event in events)) + + async def test_solve_one_turnstile_uses_fast_click_by_default(self): + calls = [] + + async def fake_start(_browser, *, fast_click=False): + calls.append(fast_click) + return {"page": object()} + + async def fake_wait(_item): + return "token-value" + + old_start = register._start_turnstile_challenge + old_wait = register._wait_turnstile_challenge + try: + register._start_turnstile_challenge = fake_start + register._wait_turnstile_challenge = fake_wait + + token = await register.solve_one_turnstile(object()) + finally: + register._start_turnstile_challenge = old_start + register._wait_turnstile_challenge = old_wait + + self.assertEqual(token, "token-value") + self.assertEqual(calls, [True]) + + async def test_wait_turnstile_logs_timeline_when_present(self): + messages = [] + register.REGISTER_LOG_MODE = "debug" + + async def fake_poll(_page, **_kwargs): + return "token-value-long" + + async def fake_put(_item, _ok): + return None + + old_poll = register._poll_turnstile_token + old_put = register._put_solver_page + old_log = register.log + try: + register._poll_turnstile_token = fake_poll + register._put_solver_page = fake_put + register.log = messages.append + item = { + "page": object(), + "trace": {}, + "timeline": { + "start": register.time.time(), + "events": [{"t": 0.1, "event": "x"}], + }, + } + + token = await register._wait_turnstile_challenge(item) + finally: + register._poll_turnstile_token = old_poll + register._put_solver_page = old_put + register.log = old_log + + self.assertEqual(token, "token-value-long") + self.assertTrue(any(message.startswith("[solver_timeline] ") for message in messages)) + + async def test_wait_turnstile_timeline_logs_solve_id_and_poll_summary(self): + messages = [] + register.REGISTER_LOG_MODE = "debug" + page = FakePage() + page.turnstile_token = "token-value-long" + timeline = register._new_solver_timeline(enabled=True) + + async def fake_put(_item, _ok): + return None + + old_attempts = register.SOLVER_POLL_ATTEMPTS + old_interval = register.SOLVER_POLL_INTERVAL_MS + old_put = register._put_solver_page + old_sleep = register.asyncio.sleep + old_log = register.log + try: + register.SOLVER_POLL_ATTEMPTS = 1 + register.SOLVER_POLL_INTERVAL_MS = 50 + + async def no_sleep(_seconds): + return None + + register.asyncio.sleep = no_sleep + register._put_solver_page = fake_put + register.log = messages.append + + token = await register._wait_turnstile_challenge({ + "page": page, + "trace": {}, + "timeline": timeline, + }) + finally: + register.SOLVER_POLL_ATTEMPTS = old_attempts + register.SOLVER_POLL_INTERVAL_MS = old_interval + register.asyncio.sleep = old_sleep + register._put_solver_page = old_put + register.log = old_log + + payload = next(message.removeprefix("[solver_timeline] ") for message in messages) + events = register.json.loads(payload) + poll_done = next(event for event in events if event["event"] == "poll_done") + + self.assertEqual(token, "token-value-long") + self.assertIn("solve_id", poll_done) + self.assertEqual(poll_done["poll_attempts"], 1) + self.assertEqual(poll_done["first_token_attempt"], 1) + self.assertGreaterEqual(poll_done["poll_read_ms_max"], 0) + + async def test_mouse_click_turnstile_retries_uses_center_clicks(self): + page = FakePage() + + old_retries = getattr(register, "SOLVER_MOUSE_CLICK_RETRIES", None) + old_interval = getattr(register, "SOLVER_MOUSE_CLICK_INTERVAL_MS", None) + old_sleep = register.asyncio.sleep + try: + register.SOLVER_MOUSE_CLICK_RETRIES = 3 + register.SOLVER_MOUSE_CLICK_INTERVAL_MS = 600 + + async def no_sleep(_seconds): + return None + + register.asyncio.sleep = no_sleep + + clicked = await register._repeat_mouse_click_turnstile(page) + finally: + register.asyncio.sleep = old_sleep + if old_retries is None: + delattr(register, "SOLVER_MOUSE_CLICK_RETRIES") + else: + register.SOLVER_MOUSE_CLICK_RETRIES = old_retries + if old_interval is None: + delattr(register, "SOLVER_MOUSE_CLICK_INTERVAL_MS") + else: + register.SOLVER_MOUSE_CLICK_INTERVAL_MS = old_interval + + self.assertTrue(clicked) + self.assertEqual(page.mouse.downs, 3) + self.assertEqual(page.mouse.ups, 3) + self.assertEqual(page.mouse.moves[-1], {"x": 160, "y": 45, "steps": 8}) + + async def test_mouse_click_turnstile_stops_when_token_appears(self): + page = FakePage() + evaluate_count = 0 + original_evaluate = page.evaluate + + async def evaluate(script): + nonlocal evaluate_count + if "cf-turnstile-response" in script: + evaluate_count += 1 + return "token-value-long" if evaluate_count > 1 else "" + return await original_evaluate(script) + + page.evaluate = evaluate + + old_retries = getattr(register, "SOLVER_MOUSE_CLICK_RETRIES", None) + old_interval = getattr(register, "SOLVER_MOUSE_CLICK_INTERVAL_MS", None) + old_sleep = register.asyncio.sleep + try: + register.SOLVER_MOUSE_CLICK_RETRIES = 3 + register.SOLVER_MOUSE_CLICK_INTERVAL_MS = 600 + + async def no_sleep(_seconds): + return None + + register.asyncio.sleep = no_sleep + + clicked = await register._repeat_mouse_click_turnstile(page) + finally: + register.asyncio.sleep = old_sleep + if old_retries is None: + delattr(register, "SOLVER_MOUSE_CLICK_RETRIES") + else: + register.SOLVER_MOUSE_CLICK_RETRIES = old_retries + if old_interval is None: + delattr(register, "SOLVER_MOUSE_CLICK_INTERVAL_MS") + else: + register.SOLVER_MOUSE_CLICK_INTERVAL_MS = old_interval + + self.assertTrue(clicked) + self.assertEqual(page.mouse.downs, 1) + self.assertEqual(evaluate_count, 2) + + async def test_prepare_signup_page_uses_configured_navigation_profile(self): + page = FakePage() + + old_wait_until = getattr(register, "PAGE_GOTO_WAIT_UNTIL", None) + old_post_wait = getattr(register, "PAGE_POST_WAIT_MS", None) + try: + register.PAGE_GOTO_WAIT_UNTIL = "domcontentloaded" + register.PAGE_POST_WAIT_MS = 500 + + await register._prepare_signup_page(page, redirect=True) + finally: + if old_wait_until is None: + delattr(register, "PAGE_GOTO_WAIT_UNTIL") + else: + register.PAGE_GOTO_WAIT_UNTIL = old_wait_until + if old_post_wait is None: + delattr(register, "PAGE_POST_WAIT_MS") + else: + register.PAGE_POST_WAIT_MS = old_post_wait + + self.assertEqual(page.goto_calls[-1]["wait_until"], "domcontentloaded") + self.assertEqual(page.waits[-1], 500) + + async def test_prepare_signup_page_can_block_static_assets(self): + page = FakePage() + + old_block_assets = getattr(register, "PAGE_BLOCK_STATIC_ASSETS", None) + try: + register.PAGE_BLOCK_STATIC_ASSETS = True + + await register._prepare_signup_page(page, redirect=True) + finally: + if old_block_assets is None: + delattr(register, "PAGE_BLOCK_STATIC_ASSETS") + else: + register.PAGE_BLOCK_STATIC_ASSETS = old_block_assets + + self.assertEqual(len(page.route_calls), 1) + self.assertEqual(page.route_calls[0]["pattern"], "**/*") + + async def test_default_solver_and_page_latency_profile_matches_accepted_optimization(self): + self.assertEqual(register.SOLVER_INITIAL_WAIT_MS, 500) + self.assertTrue(register.SOLVER_FAST_CLICK) + self.assertEqual(register.SOLVER_MOUSE_CLICK_RETRIES, 3) + self.assertEqual(register.SOLVER_MOUSE_CLICK_INTERVAL_MS, 600) + self.assertEqual(register.PAGE_GOTO_WAIT_UNTIL, "domcontentloaded") + self.assertEqual(register.PAGE_POST_WAIT_MS, 500) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime_log_analyzer.py b/tests/test_runtime_log_analyzer.py new file mode 100644 index 0000000..c2ce858 --- /dev/null +++ b/tests/test_runtime_log_analyzer.py @@ -0,0 +1,132 @@ +import unittest + +from tools.runtime_log_analyzer import ( + parse_monitor_lines, + parse_solver_timelines, + summarize_monitor_rows, + summarize_solver_timelines, +) + + +class RuntimeLogAnalyzerTests(unittest.TestCase): + def test_parses_csp_monitor_rows_and_recent_rates(self): + text = "\n".join( + [ + "[*] T:0 Q:3 phys:0 t_slot:3 q_slot:0 q_pend:0 " + "t_prod:76 t_adm:76 t_exp:0 q_sent:87 q_ret:87 q_adm:79 q_exp:0 " + "pair:76 ok:71 fail:0 rate:8.6/min #71", + "[*] T:0 Q:3 phys:0 t_slot:3 q_slot:0 q_pend:0 " + "t_prod:91 t_adm:91 t_exp:0 q_sent:101 q_ret:101 q_adm:94 q_exp:0 " + "pair:91 ok:86 fail:0 rate:8.8/min #86", + ] + ) + + rows = parse_monitor_lines(text) + summary = summarize_monitor_rows(rows) + + self.assertEqual(rows[0].kind, "csp") + self.assertEqual(summary["last_ok"], 86) + self.assertAlmostEqual(summary["last_cumulative_rate"], 8.8) + self.assertAlmostEqual(summary["recent_ok_per_min"], 9.89, places=2) + self.assertAlmostEqual(summary["recent_t_prod_per_min"], 9.89, places=2) + self.assertEqual(summary["last_q_minus_t"], 3) + self.assertEqual(summary["last_q_return_minus_t_prod"], 10) + + def test_parses_state_machine_monitor_rows(self): + text = "\n".join( + [ + "[*] slots:7/8 act:7 cpu:83% avg:76%/85 mem:4841M " + "T:4 Q:0 sent:0 got:0(0%) rate:0.0/min #0", + "[*] slots:8/8 act:8 cpu:91% avg:88%/85 mem:5012M " + "T:6 Q:5 sent:61 got:48(79%) rate:7.8/min #40", + "[*] slots:6/8 act:6 cpu:100% avg:90%/85 mem:5351M " + "T:1 Q:6 sent:68 got:68(100%) rate:8.8/min #59", + ] + ) + + rows = parse_monitor_lines(text) + summary = summarize_monitor_rows(rows) + + self.assertEqual(rows[-1].kind, "state_machine") + self.assertEqual(summary["last_ok"], 59) + self.assertAlmostEqual(summary["last_cumulative_rate"], 8.8) + self.assertEqual(summary["last_q_minus_t"], 5) + self.assertEqual(summary["last_slots"], 6) + + def test_parses_csp_v2_monitor_rows_with_admission_fields(self): + text = ( + "[*] T:1 Q:2 phys:3 p_send:1 t_slot:4 q_slot:5 q_pend:6 " + "p_batch:3.5 t_prog:1 q_inflight:4 " + "s_phys:0.50/3.00 p_phys:0.20/1.40 c_phys:0.30/2.50 " + "p_stage:0.50/0.80/0.40 c_stage:0.30/0.40/1.50 c_hot:4/1 " + "solver_goto:0.10 solver_inject:0.20 solver_initial:0.50 " + "solver_click:0.01 solver_wait:11.80 solver_reuse:0.75 solver_visible:0.00 " + "t_prod:20 t_adm:18 t_exp:1 q_sent:24 q_ret:22 q_adm:20 q_exp:0 " + "pair:17 ok:16 fail:1 rate:9.1/min #16" + ) + + rows = parse_monitor_lines(text) + summary = summarize_monitor_rows(rows) + + self.assertEqual(len(rows), 1) + self.assertEqual(rows[0].kind, "csp") + self.assertEqual(summary["last_ok"], 16) + self.assertEqual(summary["last_q_return_minus_t_prod"], 2) + self.assertEqual(summary["last_solver_wait"], 11.8) + self.assertEqual(summary["last_solver_reuse"], 0.75) + self.assertEqual(summary["last_phys"], 3) + self.assertEqual(summary["last_p_batch"], 3.5) + self.assertEqual(summary["last_t_prog"], 1) + self.assertEqual(summary["last_q_inflight"], 4) + self.assertEqual(summary["last_s_phys_hold"], 3.0) + self.assertEqual(summary["last_p_phys_wait"], 0.2) + self.assertEqual(summary["last_c_phys_hold"], 2.5) + self.assertEqual(summary["last_physical_hold_leader"], "s") + self.assertEqual(summary["last_physical_wait_leader"], "s") + self.assertEqual(summary["last_p_email_create"], 0.5) + self.assertEqual(summary["last_p_page_prepare"], 0.8) + self.assertEqual(summary["last_p_send"], 0.4) + self.assertEqual(summary["last_c_page_acquire"], 0.3) + self.assertEqual(summary["last_c_verify"], 0.4) + self.assertEqual(summary["last_c_register"], 1.5) + self.assertEqual(summary["last_c_hot_hits"], 4) + self.assertEqual(summary["last_c_hot_misses"], 1) + + def test_summarizes_solver_timeline_events(self): + text = ( + '[solver_timeline] [{"t":0.0,"event":"page_trace_after_inject",' + '"page_trace":{"created_at":0.0,"render_called_at":100.0,"render_returned_at":120.0,' + '"token_written_at":null}},' + '{"t":0.5,"event":"click_before","attempt":1,"token_len":0,' + '"dom":{"widget":{"present":true,"visible":true},' + '"turnstile_iframe_count":0,' + '"element_at_center":{"tag":"DIV","is_iframe":false}}},' + '{"t":3.0,"event":"click_after","attempt":1,"token_len":0,' + '"click_call_ms":2500.0,' + '"click_trace":{"mouse_move1_ms":500.0,"mouse_move2_ms":1500.0,' + '"mouse_down_ms":300.0,"mouse_up_ms":100.0}},' + '{"t":3.0,"event":"poll_start"},' + '{"t":4.0,"event":"poll_done","ok":true,"token_len":730,' + '"poll_attempts":2,"first_token_attempt":2,"poll_read_ms_avg":12.0,' + '"poll_read_ms_max":20.0,' + '"page_trace":{"created_at":0.0,"render_called_at":100.0,"render_returned_at":120.0,' + '"token_written_at":3900.0}}]' + ) + + timelines = parse_solver_timelines(text) + summary = summarize_solver_timelines(timelines) + + self.assertEqual(len(timelines), 1) + self.assertEqual(summary["solver_timeline_count"], 1) + self.assertEqual(summary["ok_count"], 1) + self.assertEqual(summary["avg_click_call_ms"], 2500.0) + self.assertEqual(summary["avg_click_mouse_move_ms"], 2000.0) + self.assertEqual(summary["avg_render_to_token_ms"], 3800.0) + self.assertEqual(summary["avg_token_write_to_poll_done_ms"], 100.0) + self.assertEqual(summary["avg_poll_attempts"], 2.0) + self.assertEqual(summary["center_iframe_hit_ratio"], 0.0) + self.assertEqual(summary["turnstile_iframe_seen_ratio"], 0.0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_shell_entrypoints.py b/tests/test_shell_entrypoints.py new file mode 100644 index 0000000..a9064a4 --- /dev/null +++ b/tests/test_shell_entrypoints.py @@ -0,0 +1,64 @@ +import os +import shutil +import subprocess +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[1] + + +def _entrypoint_workspace(tmp_path): + for name in ("start.sh", "auth-service.sh", "setup.sh"): + shutil.copy2(ROOT / name, tmp_path / name) + scripts = tmp_path / "scripts" + scripts.mkdir() + shutil.copy2(ROOT / "scripts" / "ensure_runtime.sh", scripts / "ensure_runtime.sh") + python = tmp_path / ".venv" / "bin" / "python" + python.parent.mkdir(parents=True) + python.write_text( + "#!/bin/sh\nprintf '%s\\n' \"$@\" > \"$ENTRY_CAPTURE\"\n", + encoding="utf-8", + ) + python.chmod(0o755) + return tmp_path + + +def _run_entry(workspace, script, *args): + capture = workspace / "capture.txt" + environment = {**os.environ, "ENTRY_CAPTURE": str(capture)} + completed = subprocess.run( + ["bash", script, *args], + cwd=workspace, + env=environment, + capture_output=True, + text=True, + ) + assert completed.returncode == 0, completed.stderr + return capture.read_text(encoding="utf-8").splitlines() + + +def test_start_dispatches_registration_and_email_as_independent_modules(tmp_path): + workspace = _entrypoint_workspace(tmp_path) + (workspace / ".env").write_text("EMAIL_MODE=tempmail\n", encoding="utf-8") + + assert _run_entry(workspace, "start.sh", "--target", "3") == [ + "-m", + "grok_register.register", + "--target", + "3", + ] + assert _run_entry(workspace, "start.sh", "--email-service", "--port", "9090") == [ + "-m", + "grok_register.email_server", + "--port", + "9090", + ] + + +def test_auth_service_keeps_the_supported_python_module_internal(tmp_path): + workspace = _entrypoint_workspace(tmp_path) + assert _run_entry(workspace, "auth-service.sh", "--debug") == [ + "-m", + "xai_enroller.service", + "--debug", + ] diff --git a/tests/test_stress.py b/tests/test_stress.py new file mode 100644 index 0000000..39d4077 --- /dev/null +++ b/tests/test_stress.py @@ -0,0 +1,411 @@ +""" +Layer 4: 压力测试 + +真实 asyncio 并发跑较长时间,检查: + - pending 泄漏 + - 库存卡死 + - claim 饥饿 + - 等待时间周期性尖峰 + - slot 守恒在长时间运行后仍成立 +""" +import asyncio +import time +import statistics +import pytest + +from grok_register.core.envelope import ResourceEnvelope +from grok_register.core.inventory import Inventory +from grok_register.core.observer import Metrics +from tests.fakes import Conservation, EventLog + + +# ═══════════════════════════════════════════ +# S/P/C 真实并发 Worker (使用 fake 服务) +# ═══════════════════════════════════════════ + +async def stress_s_worker(wid, t_gen, inventory, phys, t_slot, metrics, stop): + """压力测试 S_Worker。""" + while not stop.is_set(): + await phys.acquire() + try: + token = await t_gen.generate() + finally: + phys.release() + + if token is None: + metrics.t_discarded += 1 + continue + + metrics.t_produced += 1 + now = time.time() + try: + t_env = await ResourceEnvelope.create_with_slot( + 'T', token, t_slot, expires_at=now + 60 + ) + except asyncio.CancelledError: + metrics.t_discarded += 1 + raise + + try: + await inventory.put_t(t_env) + except Exception: + t_env.discard(reason='put_fail') + metrics.t_discarded += 1 + await asyncio.sleep(0.01) + + +async def stress_p_worker(wid, email_svc, inventory, phys, q_pend, q_slot, metrics, stop): + """压力测试 P_Worker。""" + loop = asyncio.get_event_loop() + while not stop.is_set(): + await q_pend.acquire() + _pend_released = False + try: + await phys.acquire() + try: + handle, email, password = await email_svc.create_email() + sent = True + finally: + phys.release() + + if not sent: + q_pend.release() + _pend_released = True + continue + + metrics.q_sent += 1 + + try: + code = await asyncio.wait_for(email_svc.poll_code(handle), timeout=95) + except (asyncio.TimeoutError, Exception): + code = None + + if code is None: + metrics.q_discarded += 1 + q_pend.release() + _pend_released = True + continue + + metrics.q_returned += 1 + now = time.time() + q_env = await ResourceEnvelope.create_with_slot( + 'Q', {'email': email, 'password': password, 'code': code}, + q_slot, expires_at=now + 60, + ) + + try: + await inventory.put_q(q_env) + except Exception: + q_env.discard(reason='put_fail') + metrics.q_discarded += 1 + + except asyncio.CancelledError: + if not _pend_released: + q_pend.release() + _pend_released = True + raise + except Exception: + metrics.q_discarded += 1 + finally: + if not _pend_released: + q_pend.release() + await asyncio.sleep(0.01) + + +async def stress_c_worker(wid, register_api, inventory, phys, metrics, stop, file_lock): + """压力测试 C_Worker。""" + claim_wait_times = [] + while not stop.is_set(): + t0 = time.time() + try: + async with inventory.claim_pair() as pair: + claim_wait_times.append(time.time() - t0) + t_val = pair.t.value + q_val = pair.q.value + + await phys.acquire() + try: + sso = await register_api.register( + q_val['email'], q_val['password'], q_val['code'], t_val + ) + if sso: + metrics.pair_consumed_ok += 1 + metrics.success_count += 1 + else: + metrics.pair_consumed_fail += 1 + finally: + phys.release() + + except asyncio.CancelledError: + raise + except Exception: + metrics.pair_consumed_fail += 1 + + await asyncio.sleep(0.01) + + return claim_wait_times + + +@pytest.mark.slow +class TestStress: + """压力测试: 长时间真实并发。""" + + @pytest.mark.asyncio + async def test_steady_state_30s(self): + """稳态运行 30 秒,检查守恒、无泄漏、无饥饿。""" + from tests.fakes import FakeTurnstile, FakeEmailService, FakeRegisterAPI + + t_cap, q_cap, pend_cap, phys_cap = 16, 16, 24, 8 + t_slot_sem = asyncio.Semaphore(t_cap) + q_slot_sem = asyncio.Semaphore(q_cap) + q_pend_sem = asyncio.Semaphore(pend_cap) + phys_sem = asyncio.Semaphore(phys_cap) + + ft = FakeTurnstile(delay=0.005, fail_rate=0.05) + es = FakeEmailService() + fra = FakeRegisterAPI(success_rate=0.9, delay=0.005) + + metrics = Metrics() + inv = Inventory(metrics=metrics) + cons = Conservation(t_slot_sem, q_slot_sem, q_pend_sem, inv) + stop = asyncio.Event() + fl = asyncio.Lock() + + tasks = [] + all_claim_waits = [] + + # S workers + for i in range(4): + tasks.append(asyncio.create_task( + stress_s_worker(i, ft, inv, phys_sem, t_slot_sem, metrics, stop) + )) + # P workers + for i in range(6): + tasks.append(asyncio.create_task( + stress_p_worker(i, es, inv, phys_sem, q_pend_sem, q_slot_sem, metrics, stop) + )) + # C workers + for i in range(4): + tasks.append(asyncio.create_task( + stress_c_worker(i, fra, inv, phys_sem, metrics, stop, fl) + )) + + # 监控 + 守恒检查 + async def monitor(): + snapshots = [] + while not stop.is_set(): + await asyncio.sleep(2) + try: + cons.check_t_conservation(t_cap) + cons.check_q_conservation(q_cap) + # Q_Pending: 只做软检查(非负、不超限),不做守恒检查 + # 因为 P worker 正常 acquire/release 期间 free 值会波动 + assert q_pend_sem._value >= 0, 'Q pending 为负' + assert q_pend_sem._value <= pend_cap, 'Q pending 超限' + except AssertionError as e: + snapshots.append(f'INVARIANT VIOLATION: {e}') + snapshots.append( + f't={metrics.t_produced} q_sent={metrics.q_sent} q_ret={metrics.q_returned} ' + f'pair_ok={metrics.pair_consumed_ok} inv_t={inv.t_depth} inv_q={inv.q_depth} ' + f'phys={phys_sem._value} t_slot={t_slot_sem._value} q_slot={q_slot_sem._value}' + ) + return snapshots + + mon_task = asyncio.create_task(monitor()) + + await asyncio.sleep(30) + stop.set() + for t in tasks: + t.cancel() + for t in tasks: + try: + await t + except asyncio.CancelledError: + pass + + snapshots = await mon_task + + # 最终守恒检查(所有 Worker 已停止) + cons.check_t_conservation(t_cap) + cons.check_q_conservation(q_cap) + # Q_Pending: 软检查 + assert q_pend_sem._value >= 0 and q_pend_sem._value <= pend_cap + + # 检查是否有 INVARIANT VIOLATION + violations = [s for s in snapshots if 'INVARIANT' in s] + assert not violations, f'守恒违反:\n' + '\n'.join(violations) + + print(f'\n[Stress 30s] t_produced={metrics.t_produced} ' + f'q_sent={metrics.q_sent} q_returned={metrics.q_returned} ' + f'pair_ok={metrics.pair_consumed_ok} pair_fail={metrics.pair_consumed_fail} ' + f'success={metrics.success_count}') + print(f'Snapshots ({len(snapshots)}):') + for s in snapshots[-5:]: + print(f' {s}') + + @pytest.mark.asyncio + async def test_no_claim_starvation(self): + """检查 C 不会永远等不到 pair (claim 饥饿)。""" + from tests.fakes import FakeTurnstile, FakeEmailService, FakeRegisterAPI + + t_cap, q_cap, pend_cap, phys_cap = 4, 4, 6, 2 + t_slot_sem = asyncio.Semaphore(t_cap) + q_slot_sem = asyncio.Semaphore(q_cap) + q_pend_sem = asyncio.Semaphore(pend_cap) + phys_sem = asyncio.Semaphore(phys_cap) + + ft = FakeTurnstile(delay=0.002, fail_rate=0.0) + es = FakeEmailService() + fra = FakeRegisterAPI(success_rate=1.0, delay=0.002) + + metrics = Metrics() + inv = Inventory(metrics=metrics) + stop = asyncio.Event() + fl = asyncio.Lock() + + claim_wait_samples = [] + last_success_time = time.time() + + async def c_worker(wid): + nonlocal last_success_time + while not stop.is_set(): + t0 = time.time() + try: + async with inv.claim_pair() as pair: + wait = time.time() - t0 + claim_wait_samples.append(wait) + + await phys_sem.acquire() + try: + sso = await fra.register( + pair.q.value['email'], pair.q.value['password'], + pair.q.value['code'], pair.t.value + ) + if sso: + metrics.pair_consumed_ok += 1 + metrics.success_count += 1 + last_success_time = time.time() + finally: + phys_sem.release() + except asyncio.CancelledError: + break + except Exception: + metrics.pair_consumed_fail += 1 + await asyncio.sleep(0.005) + + tasks = [] + for i in range(3): + tasks.append(asyncio.create_task(stress_s_worker(i, ft, inv, phys_sem, t_slot_sem, metrics, stop))) + for i in range(4): + tasks.append(asyncio.create_task(stress_p_worker(i, es, inv, phys_sem, q_pend_sem, q_slot_sem, metrics, stop))) + for i in range(3): + tasks.append(asyncio.create_task(c_worker(i))) + + await asyncio.sleep(15) + stop.set() + for t in tasks: + t.cancel() + for t in tasks: + try: + await t + except asyncio.CancelledError: + pass + + # 检查: claim 等待时间不应有极端尖峰 + if claim_wait_samples: + p50 = statistics.median(claim_wait_samples) + p95 = sorted(claim_wait_samples)[int(len(claim_wait_samples) * 0.95)] + p99 = sorted(claim_wait_samples)[int(len(claim_wait_samples) * 0.99)] + print(f'\n[Starvation] claim_wait p50={p50:.3f}s p95={p95:.3f}s p99={p99:.3f}s') + print(f' samples={len(claim_wait_samples)} pair_ok={metrics.pair_consumed_ok}') + # p99 不应超过 5 秒(在 fake 服务下) + assert p99 < 5.0, f'claim 饥饿: p99={p99:.3f}s 太高' + + # 检查: 最近成功消费不应太久 + time_since_last = time.time() - last_success_time + assert time_since_last < 5.0, f'太久没有成功消费: {time_since_last:.1f}s' + + @pytest.mark.asyncio + async def test_high_cancellation_churn(self): + """高取消翻转: Worker 频繁取消重启,系统仍保持守恒。""" + from tests.fakes import FakeTurnstile, FakeEmailService, FakeRegisterAPI + + t_cap, q_cap, pend_cap, phys_cap = 8, 8, 12, 4 + t_slot_sem = asyncio.Semaphore(t_cap) + q_slot_sem = asyncio.Semaphore(q_cap) + q_pend_sem = asyncio.Semaphore(pend_cap) + phys_sem = asyncio.Semaphore(phys_cap) + + ft = FakeTurnstile(delay=0.003, fail_rate=0.1) + es = FakeEmailService() + fra = FakeRegisterAPI(success_rate=0.85, delay=0.003) + + metrics = Metrics() + inv = Inventory(metrics=metrics) + cons = Conservation(t_slot_sem, q_slot_sem, q_pend_sem, inv) + stop = asyncio.Event() + fl = asyncio.Lock() + + tasks = [] + violations = [] + + def spawn_worker(kind, idx): + if kind == 'S': + return asyncio.create_task(stress_s_worker(idx, ft, inv, phys_sem, t_slot_sem, metrics, stop)) + elif kind == 'P': + return asyncio.create_task(stress_p_worker(idx, es, inv, phys_sem, q_pend_sem, q_slot_sem, metrics, stop)) + else: + return asyncio.create_task(stress_c_worker(idx, fra, inv, phys_sem, metrics, stop, fl)) + + # 初始 workers + for i in range(3): + tasks.append(spawn_worker('S', i)) + tasks.append(spawn_worker('P', i)) + tasks.append(spawn_worker('C', i)) + + async def churn(): + import random + while not stop.is_set(): + await asyncio.sleep(0.3) + # 随机取消一个,重启一个 + if tasks: + idx = random.randint(0, len(tasks) - 1) + t = tasks[idx] + if not t.done(): + t.cancel() + tasks[idx] = spawn_worker( + random.choice(['S', 'P', 'C']), + random.randint(0, 9) + ) + # 定期守恒检查(软检查) + try: + cons.check_t_conservation(t_cap) + cons.check_q_conservation(q_cap) + assert q_pend_sem._value >= 0 and q_pend_sem._value <= pend_cap + except AssertionError as e: + violations.append(str(e)) + + churn_task = asyncio.create_task(churn()) + tasks.append(churn_task) + + await asyncio.sleep(15) + stop.set() + for t in tasks: + t.cancel() + for t in tasks: + try: + await t + except asyncio.CancelledError: + pass + + # 最终守恒检查 + try: + cons.check_t_conservation(t_cap) + cons.check_q_conservation(q_cap) + assert q_pend_sem._value >= 0 and q_pend_sem._value <= pend_cap + except AssertionError as e: + violations.append(f'FINAL: {e}') + + assert not violations, f'守恒违反:\n' + '\n'.join(violations[:5]) + print(f'\n[High Churn] t_produced={metrics.t_produced} ' + f'pair_ok={metrics.pair_consumed_ok} success={metrics.success_count}') diff --git a/tests/test_xai_auth_pipeline.py b/tests/test_xai_auth_pipeline.py new file mode 100644 index 0000000..9a8eaa1 --- /dev/null +++ b/tests/test_xai_auth_pipeline.py @@ -0,0 +1,298 @@ +import asyncio +import threading + +from xai_enroller.auth_pipeline import AuthPipeline +from xai_enroller.ledger import Ledger +from xai_enroller.models import ( + AuthorizationResult, + AuthorizationStatus, + DeviceFlow, + JobStatus, + OAuthCredential, + PipelineState, + SinkReceipt, + SourceRecord, +) + + +class EmptySource: + async def records(self): + if False: + yield None + + async def close(self): + return None + + +class NoopExecutor: + async def close(self): + return None + + +def _pipeline(tmp_path): + return AuthPipeline( + source=EmptySource(), + protocol=object(), + executor=NoopExecutor(), + sink=object(), + ledger=Ledger(tmp_path / "ledger.db", b"salt"), + ) + + +def test_cancel_active_is_idempotent_while_cleanup_is_in_progress(tmp_path): + async def scenario(): + pipeline = _pipeline(tmp_path) + + async def wait_forever(): + await asyncio.Future() + + task = asyncio.create_task(wait_forever()) + pipeline._authorization_task = task + pipeline._authorization_cancellable = True + assert await pipeline.cancel_active() + assert not await pipeline.cancel_active() + await asyncio.gather(task, return_exceptions=True) + + asyncio.run(scenario()) + + +class SingleSource: + async def records(self): + yield SourceRecord("stock", "opaque") + + async def close(self): + return None + + +class ImmediateProtocol: + async def start_device_flow(self): + return DeviceFlow( + "device", + "code", + "https://accounts.x.ai/oauth2/device", + 600, + 0, + ) + + async def poll_token(self, **_kwargs): + return OAuthCredential( + "access", + "refresh", + None, + "Bearer", + 3600, + "later", + "now", + "subject", + "https://auth.x.ai/oauth2/token", + ) + + +class ImmediateExecutor: + def __init__(self): + self.calls = 0 + + async def start(self): + return None + + async def close(self): + return None + + async def confirm(self, *_args): + self.calls += 1 + return AuthorizationResult(AuthorizationStatus.AUTHORIZED, "confirmed") + + +class BlockingThreadSink: + def __init__(self): + self.entered = threading.Event() + self.release = threading.Event() + self.pipeline = None + self.calls = 0 + + def _store(self): + self.calls += 1 + self.entered.set() + self.release.wait(2) + return SinkReceipt("opaque") + + async def store(self, _credential): + receipt = await asyncio.to_thread(self._store) + self.pipeline.request_stop() + return receipt + + +def test_sink_and_ledger_commit_cannot_be_cancelled_mid_commit(tmp_path): + async def scenario(): + executor = ImmediateExecutor() + sink = BlockingThreadSink() + ledger = Ledger(tmp_path / "ledger.db", b"salt") + events = [] + pipeline = AuthPipeline( + source=SingleSource(), + protocol=ImmediateProtocol(), + executor=executor, + sink=sink, + ledger=ledger, + timeout=2, + event_callback=lambda kind, data: events.append((kind, data)), + ) + sink.pipeline = pipeline + run = asyncio.create_task(pipeline.run()) + assert await asyncio.to_thread(sink.entered.wait, 1) + assert not await pipeline.cancel_active() + sink.release.set() + await asyncio.wait_for(run, 2) + assert executor.calls == 1 + assert sink.calls == 1 + assert ledger.aggregate_counts()["imported_unique"] == 1 + started = next(data for kind, data in events if kind == "authorization_started") + imported = next( + data + for kind, data in events + if kind == "result" and data["status"] == "imported" + ) + assert imported["task_number"] == started["task_number"] + + asyncio.run(scenario()) + + +def test_pending_total_uses_current_snapshot_set_difference(tmp_path): + class SnapshotSource(EmptySource): + snapshot_fingerprints = frozenset({"a", "b", "c"}) + + ledger = Ledger(tmp_path / "ledger.db", b"salt") + imported = ledger.start_fingerprint("b") + ledger.finish(imported, JobStatus.IMPORTED, "imported", "receipt-b") + outside = ledger.start_fingerprint("outside") + ledger.finish(outside, JobStatus.IMPORTED, "imported", "receipt-outside") + pipeline = AuthPipeline( + source=SnapshotSource(), + protocol=object(), + executor=NoopExecutor(), + sink=object(), + ledger=ledger, + ) + + assert pipeline.status()["pending_total"] == 2 + assert pipeline.status()["five_minute_imports_per_minute"] is None + assert pipeline.status()["lifetime_imports_per_minute"] is None + + +def test_pending_total_is_unknown_before_first_valid_snapshot(tmp_path): + pipeline = _pipeline(tmp_path) + + assert pipeline.status()["pending_total"] is None + + +def test_rate_limit_gate_grows_cooldown_on_consecutive_trips(): + class Clock: + def __init__(self): + self.now = 0.0 + + def __call__(self): + return self.now + + async def scenario(): + from xai_enroller.auth_pipeline import GlobalRateLimitGate + + clock = Clock() + gate = GlobalRateLimitGate(clock=clock, base_cooldown=60.0) + assert await gate.rate_limited() + assert gate.COOLDOWN_SECONDS == 60.0 + assert gate.snapshot()["cooldown_seconds"] == 60.0 + + resume = asyncio.Event() + resume.set() + clock.now = 60.0 + probe = await gate.wait_for_permission(resume) + assert probe is not None + assert await gate.rate_limited(probe) + assert gate.COOLDOWN_SECONDS == 90.0 + assert gate.snapshot()["cooldown_remaining_seconds"] == 90.0 + + clock.now = 150.0 + probe = await gate.wait_for_permission(resume) + elapsed = await gate.authorized(probe) + assert elapsed == 150.0 + assert gate.COOLDOWN_SECONDS == 60.0 + assert gate.snapshot()["rate_limit_trips"] == 0 + + asyncio.run(scenario()) + + +def test_minimum_start_interval_raises_after_rate_limit_and_decays(): + from xai_enroller.auth_pipeline import MinimumStartInterval + + def approx(value, rel=1e-6): + class Approx: + def __eq__(self, other): + return abs(float(other) - float(value)) <= rel * max(1.0, abs(value)) + + def __repr__(self): + return f"approx({value})" + + return Approx() + + interval = MinimumStartInterval(10.0) + assert interval.seconds == 10.0 + raised = interval.note_rate_limited() + assert raised == 18.0 + raised = interval.note_rate_limited() + assert raised == approx(18.0 * 1.55) + for _ in range(MinimumStartInterval.SUCCESS_STREAK_TO_DECAY): + interval.note_success() + assert interval.seconds < raised + # Keep decaying until the base floor is restored. + for _ in range(40): + interval.note_success() + assert interval.seconds == 10.0 + + +def test_queued_source_imported_while_paused_is_not_authorized(tmp_path): + async def scenario(): + executor = ImmediateExecutor() + ledger = Ledger(tmp_path / "ledger.db", b"salt") + pipeline = AuthPipeline( + source=EmptySource(), + protocol=ImmediateProtocol(), + executor=executor, + sink=object(), + ledger=ledger, + timeout=2, + ) + source = SourceRecord("stock", "opaque") + fingerprint = ledger.fingerprint(source.source_id) + + assert await pipeline.admit(source) + pipeline.pause() + await pipeline.rate_gate.rate_limited() + run = asyncio.create_task(pipeline.run()) + try: + async with asyncio.timeout(1): + while pipeline._authorization_task is None: + await asyncio.sleep(0) + + external_job_id = ledger.start_fingerprint(fingerprint) + ledger.finish( + external_job_id, + JobStatus.IMPORTED, + "imported", + "external-receipt", + ) + pipeline.resume() + + async with asyncio.timeout(1): + while ( + executor.calls == 0 + and pipeline._states.get(fingerprint) is not PipelineState.IMPORTED + ): + await asyncio.sleep(0) + + assert executor.calls == 0 + assert pipeline._states[fingerprint] is PipelineState.IMPORTED + finally: + pipeline.request_stop() + await asyncio.wait_for(run, 2) + + asyncio.run(scenario()) diff --git a/tests/test_xai_auth_service.py b/tests/test_xai_auth_service.py new file mode 100644 index 0000000..353ae9c --- /dev/null +++ b/tests/test_xai_auth_service.py @@ -0,0 +1,391 @@ +import asyncio +import io +import os +import subprocess +import sys +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from xai_enroller.models import JobStatus +from xai_enroller.service import ( + AuthServiceSettings, + AuthPipelineRunner, + EventTerminal, + InteractiveCommandPrompt, + resolve_auth_log_mode, +) +from xai_enroller.ledger import Ledger +from xai_enroller.inventory import InventoryError + + +def test_auth_log_mode_prefers_cli_and_rejects_invalid_values(): + assert resolve_auth_log_mode([], {}) == "user" + assert resolve_auth_log_mode( + [], {"XAI_AUTH_SERVICE_LOG_MODE": "debug"} + ) == "debug" + assert resolve_auth_log_mode( + ["--debug"], {"XAI_AUTH_SERVICE_LOG_MODE": "user"} + ) == "debug" + with pytest.raises(ValueError, match="XAI_AUTH_SERVICE_LOG_MODE"): + resolve_auth_log_mode([], {"XAI_AUTH_SERVICE_LOG_MODE": "verbose"}) + + +def test_user_terminal_reports_progress_without_internal_or_secret_fields(): + messages = [] + terminal = EventTerminal(mode="user", output=messages.append) + terminal.emit( + ( + "startup", + { + "available": 7, + "claimed": 3, + "destination": "authenticated/", + "source_kind": "local", + "ssh_host": "do-not-print.example", + }, + ) + ) + terminal.emit( + ( + "authorization_started", + { + "task_number": 3, + "attempt_number": 1, + "pending_total": 41, + "source_queue": 64, + "source_id": "secret@example.test", + }, + ) + ) + terminal.emit( + ( + "result", + { + "status": "imported", + "reason": "imported", + "task_number": 3, + "five_minute_imports_per_minute": 4.25, + "lifetime_imports_per_minute": 1.75, + "imported_unique": 120, + "available": 117, + "source_id": "secret@example.test", + }, + ) + ) + + assert messages == [ + "[✓] 本地认证服务已启动 | 来源 本机 | 输出 authenticated/ | 待处理 — | 可用 7", + "[→] 开始认证 #3 | 待处理 41", + "[✓] 认证成功 #3 | 运行平均 1.75/分 | 累计 120 | 可用 117", + ] + rendered = "\n".join(messages) + assert "source_queue" not in rendered + assert "secret@example.test" not in rendered + assert "do-not-print.example" not in rendered + + +def test_user_status_distinguishes_unknown_runtime_rate_from_zero_rate(): + base = { + "state": "running", + "pending_total": None, + "active_stage": "idle", + "imported_unique": 0, + "available": 0, + "claimed": 0, + "cooldown": False, + } + messages = [] + terminal = EventTerminal(mode="user", output=messages.append) + terminal.emit(("status", {**base, "lifetime_imports_per_minute": None})) + terminal.emit(("status", {**base, "lifetime_imports_per_minute": 0.0})) + + assert "运行平均 —" in messages[0] + assert "运行平均 0.00/分" in messages[1] + assert "source_queue" not in messages[0] + + +def test_user_terminal_omits_missing_task_number_before_authorization_starts(): + messages = [] + terminal = EventTerminal(mode="user", output=messages.append) + + terminal.emit( + ( + "result", + { + "status": "failed", + "reason": "device_flow_failed", + "task_number": None, + }, + ) + ) + + assert messages == ["[✗] 认证未完成 | 暂时失败,将自动重试"] + assert "None" not in messages[0] + + +def test_debug_terminal_keeps_aggregate_diagnostics_and_sanitizes_unknown_events(): + messages = [] + terminal = EventTerminal(mode="debug", output=messages.append) + terminal.emit( + ( + "status", + { + "state": "running", + "source_queue": 2, + "prepared_queue": 1, + "completion_queue": 0, + "active_stage": "authorization", + "retry_waiting": 1, + "next_retry_seconds": 5.0, + "authorization_starts": 3, + "cooldown": False, + "cooldown_remaining_seconds": 0.0, + "probe_in_flight": False, + "min_authorization_interval_seconds": 10.0, + "pacing_remaining_seconds": 2.0, + "imported_unique": 2, + "attempted_unique": 3, + "rate_limited": 1, + "five_minute_imports_per_minute": 2.0, + "lifetime_imports_per_minute": 1.0, + "available": 2, + "claiming": 0, + "claimed": 0, + }, + ) + ) + terminal.emit( + ( + "future_event", + {"reason": "internal_error", "token": "do-not-print-token"}, + ) + ) + + assert "queues=2/1/0" in messages[0] + assert messages[1] == "• debug event=future_event reason=internal_error" + assert "do-not-print-token" not in "\n".join(messages) + + +def test_terminal_output_failure_does_not_escape(): + def broken_output(_message): + raise OSError("closed terminal") + + terminal = EventTerminal(mode="user", output=broken_output) + terminal.emit(("service_stopped", {})) + + +def test_interactive_prompt_restores_partially_typed_take_command_after_events(): + output = io.StringIO() + commands = [] + prompt = InteractiveCommandPrompt( + output=output, + interactive=True, + prompt="认证> ", + ) + prompt.start(commands.append) + prompt.feed("take 12") + + prompt.write_event("[↻] 发现新账号 3") + + assert output.getvalue().endswith("[↻] 发现新账号 3\n认证> take 12") + + prompt.feed("\r") + + assert commands == ["take 12"] + assert output.getvalue().endswith("\n认证> ") + + +def test_auth_service_configuration_errors_are_actionable_without_traceback(tmp_path): + environment = os.environ.copy() + environment.pop("XAI_AUTH_SERVICE_SSH_HOST", None) + environment.pop("XAI_AUTH_SERVICE_LOG_MODE", None) + environment["XAI_AUTH_SERVICE_SOURCE"] = "ssh" + missing = subprocess.run( + [sys.executable, "-m", "xai_enroller.service"], + env=environment, + capture_output=True, + text=True, + ) + assert missing.returncode == 2 + assert "XAI_AUTH_SERVICE_SSH_HOST" in missing.stderr + assert "docs/guides/auth-service.md" in missing.stderr + assert "Traceback" not in missing.stderr + + environment["XAI_AUTH_SERVICE_LOG_MODE"] = "verbose" + invalid = subprocess.run( + [sys.executable, "-m", "xai_enroller.service"], + env=environment, + capture_output=True, + text=True, + ) + assert invalid.returncode == 2 + assert "XAI_AUTH_SERVICE_LOG_MODE" in invalid.stderr + assert "Traceback" not in invalid.stderr + + environment["XAI_AUTH_SERVICE_LOG_MODE"] = "user" + environment.pop("XAI_AUTH_SERVICE_SOURCE", None) + environment["XAI_AUTH_SERVICE_SSH_HOST"] = "user@example.test" + environment["XAI_ENROLLER_LOCAL_AUTH_DIR"] = str(tmp_path / "auth") + environment["XAI_ENROLLER_TIMEOUT_SEC"] = "not-a-number" + invalid_enroller_setting = subprocess.run( + [sys.executable, "-m", "xai_enroller.service"], + env=environment, + capture_output=True, + text=True, + ) + assert invalid_enroller_setting.returncode == 2 + assert "XAI_ENROLLER_TIMEOUT_SEC" in invalid_enroller_setting.stderr + assert "docs/guides/auth-service.md" in invalid_enroller_setting.stderr + assert "Traceback" not in invalid_enroller_setting.stderr + + +def test_auth_service_startup_failure_is_sanitized_without_traceback(tmp_path): + blocked_destination = tmp_path / "private-output-path" + blocked_destination.write_text("not a directory", encoding="utf-8") + environment = os.environ.copy() + environment["XAI_AUTH_SERVICE_SSH_HOST"] = "user@example.test" + environment["XAI_ENROLLER_LOCAL_AUTH_DIR"] = str(blocked_destination) + + failed = subprocess.run( + [sys.executable, "-m", "xai_enroller.service"], + env=environment, + capture_output=True, + text=True, + ) + + assert failed.returncode == 1 + assert "认证服务异常终止" in failed.stderr + assert "bash auth-service.sh --debug" in failed.stderr + assert "Traceback" not in failed.stderr + assert str(blocked_destination) not in failed.stderr + + +def test_ledger_recognizes_previously_imported_sources(tmp_path): + ledger = Ledger(tmp_path / "ledger.db", b"test-salt") + imported = ledger.start("done@example.test") + ledger.finish(imported, JobStatus.IMPORTED, "imported") + failed = ledger.start("retry@example.test") + ledger.finish(failed, JobStatus.SINK_FAILED, "sink_failed") + + assert ledger.has_imported("done@example.test") is True + assert ledger.has_imported("retry@example.test") is False + + +def test_auth_service_settings_defaults_local_and_preserves_ssh_auto_detection(tmp_path): + local = AuthServiceSettings.from_environ( + {"XAI_AUTH_SERVICE_REGISTER_ROOT": str(tmp_path)} + ) + remote = AuthServiceSettings.from_environ( + { + "XAI_AUTH_SERVICE_SSH_HOST": "ubuntu@example.test", + "XAI_AUTH_SERVICE_SYNC_SEC": "45", + } + ) + + assert local.source_kind == "local" + assert local.ssh_host is None + assert local.register_root == str(tmp_path) + assert remote.source_kind == "ssh" + assert remote.ssh_host == "ubuntu@example.test" + assert remote.sync_seconds == 45 + assert remote.remote_root == "/opt/grok-free-register" + + +def test_auth_service_settings_explicit_source_overrides_auto_detection(): + local = AuthServiceSettings.from_environ( + { + "XAI_AUTH_SERVICE_SOURCE": "local", + "XAI_AUTH_SERVICE_SSH_HOST": "ignored@example.test", + } + ) + assert local.source_kind == "local" + + with pytest.raises(ValueError, match="XAI_AUTH_SERVICE_SSH_HOST"): + AuthServiceSettings.from_environ({"XAI_AUTH_SERVICE_SOURCE": "ssh"}) + with pytest.raises(ValueError, match="XAI_AUTH_SERVICE_SOURCE"): + AuthServiceSettings.from_environ({"XAI_AUTH_SERVICE_SOURCE": "unknown"}) + + +def test_pipeline_runner_takes_a_credential_batch_and_reports_inventory(): + class InventoryLedger: + def inventory_counts(self): + return {"available": 7, "claiming": 0, "claimed": 3} + + class Pipeline: + ledger = InventoryLedger() + + def status(self): + return {"state": "running"} + + class Inventory: + def take(self, count): + assert count == 3 + return SimpleNamespace( + batch_id="batch-1", + directory=Path("/tmp/claimed/batch-1"), + moved=3, + note="", + ) + + async def scenario(): + events = [] + runner = AuthPipelineRunner(Pipeline(), events.append, inventory=Inventory()) + assert await runner.handle_command("take 3") is True + assert await runner.handle_command("s") is True + return events + + assert asyncio.run(scenario()) == [ + ( + "inventory_taken", + { + "batch_id": "batch-1", + "directory": "/tmp/claimed/batch-1", + "moved": 3, + "available": 7, + "claiming": 0, + "claimed": 3, + }, + ), + ( + "status", + { + "state": "running", + "available": 7, + "claiming": 0, + "claimed": 3, + }, + ), + ] + + +def test_pipeline_runner_reports_inventory_failure_without_stopping(): + class Ledger: + def inventory_counts(self): + return {"available": 1, "claiming": 1, "claimed": 0} + + class Pipeline: + ledger = Ledger() + + class Inventory: + def take(self, _count): + raise InventoryError("credential file is missing") + + async def scenario(): + events = [] + runner = AuthPipelineRunner(Pipeline(), events.append, inventory=Inventory()) + assert await runner.handle_command("take 1") is True + return events + + assert asyncio.run(scenario()) == [ + ( + "inventory_error", + { + "reason": "credential file is missing", + "available": 1, + "claiming": 1, + "claimed": 0, + }, + ) + ] diff --git a/tests/test_xai_credential_inventory.py b/tests/test_xai_credential_inventory.py new file mode 100644 index 0000000..6d0d271 --- /dev/null +++ b/tests/test_xai_credential_inventory.py @@ -0,0 +1,58 @@ +from xai_enroller.inventory import CredentialInventory +from xai_enroller.ledger import Ledger +from xai_enroller.models import JobStatus + + +def import_credential(ledger, source, receipt): + job_id = ledger.start(source) + ledger.finish(job_id, JobStatus.IMPORTED, "imported", receipt) + return job_id + + +def test_take_moves_available_credentials_without_changing_imported_jobs(tmp_path): + ledger = Ledger(tmp_path / "ledger.db", b"salt") + available = tmp_path / "authenticated" + claimed = tmp_path / "claimed" + available.mkdir() + first_job = import_credential(ledger, "first", "xai-first") + second_job = import_credential(ledger, "second", "xai-second") + (available / "xai-first.json").write_text("{}\n", encoding="utf-8") + (available / "xai-second.json").write_text("{}\n", encoding="utf-8") + + batch = CredentialInventory(ledger, available, claimed).take(2) + + assert batch.moved == 2 + assert batch.note == "" + assert not list(available.glob("*.json")) + assert sorted(path.name for path in batch.directory.glob("*.json")) == [ + "xai-first.json", + "xai-second.json", + ] + assert ledger.inventory_counts() == { + "available": 0, + "claiming": 0, + "claimed": 2, + } + assert ledger.get(first_job)["status"] == JobStatus.IMPORTED.value + assert ledger.get(second_job)["status"] == JobStatus.IMPORTED.value + + +def test_recover_finishes_a_batch_interrupted_before_file_move(tmp_path): + ledger = Ledger(tmp_path / "ledger.db", b"salt") + available = tmp_path / "authenticated" + claimed = tmp_path / "claimed" + available.mkdir() + import_credential(ledger, "source", "xai-recover") + (available / "xai-recover.json").write_text("{}\n", encoding="utf-8") + assert ledger.claim_available(1, "batch-recover") == ["xai-recover"] + + recovered = CredentialInventory(ledger, available, claimed).recover() + + assert recovered == 1 + assert not (available / "xai-recover.json").exists() + assert (claimed / "batch-recover" / "xai-recover.json").exists() + assert ledger.inventory_counts() == { + "available": 0, + "claiming": 0, + "claimed": 1, + } diff --git a/tests/test_xai_enroller_config.py b/tests/test_xai_enroller_config.py new file mode 100644 index 0000000..85824a3 --- /dev/null +++ b/tests/test_xai_enroller_config.py @@ -0,0 +1,125 @@ +import os +import sqlite3 +import stat + +import pytest + +from xai_enroller.config import Settings +from xai_enroller.sources import FileSourceAdapter, SQLiteSourceAdapter + + +def base_env(tmp_path, **overrides): + values = { + "XAI_ENROLLER_SOURCE_KIND": "file", + "XAI_ENROLLER_SOURCE_FILE": str(tmp_path / "sessions.tsv"), + "XAI_ENROLLER_SOURCE_SALT": "test-salt", + "XAI_ENROLLER_LEDGER_PATH": str(tmp_path / "ledger.db"), + } + values.update(overrides) + return values + + +def test_settings_rejects_public_http_cpa_sink(tmp_path): + with pytest.raises(ValueError, match="HTTPS"): + Settings.from_environ( + base_env( + tmp_path, + XAI_ENROLLER_SINK="cpa", + XAI_ENROLLER_CPA_BASE_URL="http://example.test", + XAI_ENROLLER_CPA_MANAGEMENT_SECRET="secret", + ) + ) + + +def test_settings_defaults_are_bounded_and_redacted(tmp_path): + settings = Settings.from_environ(base_env(tmp_path)) + assert settings.target == 1 + assert settings.concurrency == 1 + assert settings.retry_attempts == 0 + assert settings.executor == "http" + assert settings.poll_interval == 5.0 + redacted = settings.redacted_dict() + assert "test-salt" not in repr(redacted) + assert "CPA_MANAGEMENT_SECRET" not in repr(redacted) + assert "source_salt" not in redacted + + +def test_settings_accepts_bounded_retry_attempts(tmp_path): + settings = Settings.from_environ( + base_env(tmp_path, XAI_ENROLLER_RETRY_ATTEMPTS="2") + ) + + assert settings.retry_attempts == 2 + + +def test_settings_accepts_remote_source_without_a_local_source_file(tmp_path): + settings = Settings.from_environ( + base_env( + tmp_path, + XAI_ENROLLER_SOURCE_KIND="remote", + XAI_ENROLLER_SOURCE_FILE="", + ) + ) + + assert settings.source_kind == "remote" + assert settings.source_file is None + assert settings.source_db is None + + +@pytest.mark.parametrize( + ("key", "value", "message"), + [ + ("XAI_ENROLLER_SOURCE_KIND", "other", "source"), + ("XAI_ENROLLER_AUTH_EXECUTOR", "other", "executor"), + ("XAI_ENROLLER_CONCURRENCY", "0", "concurrency"), + ("XAI_ENROLLER_POLL_SEC", "0", "poll"), + ("XAI_ENROLLER_RETRY_ATTEMPTS", "-1", "retry"), + ("XAI_ENROLLER_RETRY_ATTEMPTS", "4", "retry"), + ("XAI_ENROLLER_TARGET", "0", "target"), + ("XAI_ENROLLER_TARGET", "101", "target"), + ], +) +def test_settings_reject_invalid_bounds(tmp_path, key, value, message): + with pytest.raises(ValueError, match=message): + Settings.from_environ(base_env(tmp_path, **{key: value})) + + +def test_file_source_requires_private_regular_file(tmp_path): + path = tmp_path / "sessions.tsv" + path.write_text("one\tsecret\n", encoding="utf-8") + path.chmod(stat.S_IRUSR | stat.S_IWUSR | stat.S_IRGRP) + with pytest.raises(ValueError, match="0600"): + list(FileSourceAdapter(path).records()) + + +def test_file_source_suppresses_duplicates_without_leaking_tokens(tmp_path): + path = tmp_path / "sessions.tsv" + path.write_text("one\tsecret-a\none\tsecret-b\n", encoding="utf-8") + path.chmod(0o600) + records = list(FileSourceAdapter(path).records()) + assert len(records) == 1 + assert records[0].source_id == "one" + assert "secret" not in repr(records) + + +def test_sqlite_source_uses_fixed_read_only_query_and_hmac_ids(tmp_path): + path = tmp_path / "accounts.db" + connection = sqlite3.connect(path) + connection.execute( + "CREATE TABLE accounts (token TEXT, status TEXT, deleted_at TEXT, updated_at INTEGER)" + ) + connection.executemany( + "INSERT INTO accounts VALUES (?, ?, ?, ?)", + [ + ("old", "active", None, 1), + ("new", "active", None, 2), + ("deleted", "active", "2026-01-01", 3), + ("inactive", "inactive", None, 4), + ], + ) + connection.commit() + connection.close() + records = list(SQLiteSourceAdapter(path, b"salt").records()) + assert [record.sso_token for record in records] == ["new", "old"] + assert records[0].source_id != records[1].source_id + assert "new" not in repr(records) diff --git a/tests/test_xai_enroller_coordinator.py b/tests/test_xai_enroller_coordinator.py new file mode 100644 index 0000000..4a0883d --- /dev/null +++ b/tests/test_xai_enroller_coordinator.py @@ -0,0 +1,293 @@ +import asyncio +import sqlite3 + +import pytest + +from xai_enroller.coordinator import EnrollmentCoordinator +from xai_enroller.models import ( + AuthorizationResult, + AuthorizationStatus, + DeviceFlow, + JobStatus, + OAuthCredential, + SinkReceipt, + SourceRecord, +) +from xai_enroller.sources import FileSourceAdapter, SQLiteSourceAdapter + + +class Source: + async def records(self): + for index in range(3): + yield SourceRecord(f"source-{index}", f"token-{index}") + + +class Protocol: + def __init__(self): + self.polls = 0 + + async def start_device_flow(self): + return DeviceFlow("device", "code", "https://accounts.x.ai/device", 60, 0) + + async def poll_token(self, **kwargs): + self.polls += 1 + return OAuthCredential( + "access", + "refresh", + None, + "Bearer", + 3600, + "2026-07-11T00:00:00Z", + "2026-07-11T00:00:00Z", + "subject", + "https://auth.x.ai/oauth2/token", + ) + + +class Executor: + active = 0 + max_active = 0 + + async def confirm(self, source, flow): + type(self).active += 1 + type(self).max_active = max(type(self).max_active, type(self).active) + await asyncio.sleep(0) + type(self).active -= 1 + if source.source_id == "source-1": + return AuthorizationResult(AuthorizationStatus.NEEDS_BROWSER, "challenge") + return AuthorizationResult(AuthorizationStatus.AUTHORIZED, "confirmed") + + +class Sink: + async def store(self, credential): + if credential.subject == "subject": + return SinkReceipt("receipt") + + +class FlakyProtocol(Protocol): + def __init__(self): + super().__init__() + self.calls = 0 + + async def start_device_flow(self): + self.calls += 1 + if self.calls == 1: + raise OSError("temporary transport") + return await super().start_device_flow() + + +class InvalidProtocol(Protocol): + def __init__(self): + super().__init__() + self.calls = 0 + + async def start_device_flow(self): + self.calls += 1 + raise ValueError("invalid device response") + + +class FailingExecutor(Executor): + async def confirm(self, source, flow): + raise OSError("confirmation transport") + + +class SlowDeviceFlowProtocol(Protocol): + async def start_device_flow(self): + await asyncio.sleep(0.04) + return await super().start_device_flow() + + +class SlowExecutor(Executor): + async def confirm(self, source, flow): + await asyncio.sleep(0.04) + return await super().confirm(source, flow) + + +def test_coordinator_limits_concurrency_and_maps_terminal_results(tmp_path): + protocol = Protocol() + coordinator = EnrollmentCoordinator( + source=Source(), + protocol=protocol, + executor=Executor(), + sink=Sink(), + ledger_path=tmp_path / "ledger.db", + ledger_salt=b"salt", + concurrency=2, + timeout=5, + ) + results = asyncio.run(coordinator.run(target=3)) + assert [result.status for result in results] == [ + JobStatus.IMPORTED, + JobStatus.NEEDS_BROWSER, + JobStatus.IMPORTED, + ] + assert Executor.max_active == 2 + assert protocol.polls == 2 + + +def test_coordinator_accepts_explicit_synced_records(tmp_path): + coordinator = EnrollmentCoordinator( + source=Source(), + protocol=Protocol(), + executor=Executor(), + sink=Sink(), + ledger_path=tmp_path / "ledger.db", + ledger_salt=b"salt", + concurrency=1, + timeout=5, + ) + + results = asyncio.run( + coordinator.run_records([SourceRecord("synced-account", "synced-sso")]) + ) + + assert [result.source_id for result in results] == ["synced-account"] + assert results[0].status is JobStatus.IMPORTED + + +def test_coordinator_emits_device_flow_after_the_flow_is_created(tmp_path): + events = [] + coordinator = EnrollmentCoordinator( + source=Source(), + protocol=Protocol(), + executor=Executor(), + sink=Sink(), + ledger_path=tmp_path / "ledger.db", + ledger_salt=b"salt", + concurrency=1, + timeout=5, + event_callback=lambda kind, data: events.append((kind, data)), + ) + + asyncio.run(coordinator.run_records([SourceRecord("synced-account", "synced-sso")])) + + assert events == [ + ( + "device_flow", + { + "source_id": "synced-account", + "user_code": "code", + "verification_url": "https://accounts.x.ai/device", + }, + ) + ] + + +def test_coordinator_without_sink_never_reports_imported(tmp_path): + coordinator = EnrollmentCoordinator( + source=Source(), + protocol=Protocol(), + executor=Executor(), + sink=None, + ledger_path=tmp_path / "ledger.db", + ledger_salt=b"salt", + concurrency=1, + timeout=5, + ) + results = asyncio.run(coordinator.run(target=1)) + assert results[0].status is JobStatus.SINK_FAILED + + +def test_coordinator_retries_only_before_device_flow_creation(tmp_path): + protocol = FlakyProtocol() + coordinator = EnrollmentCoordinator( + source=Source(), + protocol=protocol, + executor=Executor(), + sink=Sink(), + ledger_path=tmp_path / "ledger.db", + ledger_salt=b"salt", + concurrency=1, + timeout=5, + retry_attempts=1, + ) + results = asyncio.run(coordinator.run(target=1)) + assert results[0].status is JobStatus.IMPORTED + assert protocol.calls == 2 + + +def test_coordinator_does_not_retry_after_device_flow_creation(tmp_path): + protocol = Protocol() + coordinator = EnrollmentCoordinator( + source=Source(), + protocol=protocol, + executor=FailingExecutor(), + sink=Sink(), + ledger_path=tmp_path / "ledger.db", + ledger_salt=b"salt", + concurrency=1, + timeout=5, + retry_attempts=3, + ) + results = asyncio.run(coordinator.run(target=1)) + assert results[0].status is JobStatus.TRANSPORT_FAILED + assert protocol.polls == 0 + + +def test_coordinator_does_not_retry_non_transport_device_flow_error(tmp_path): + protocol = InvalidProtocol() + coordinator = EnrollmentCoordinator( + source=Source(), + protocol=protocol, + executor=Executor(), + sink=Sink(), + ledger_path=tmp_path / "ledger.db", + ledger_salt=b"salt", + concurrency=1, + timeout=5, + retry_attempts=3, + ) + results = asyncio.run(coordinator.run(target=1)) + assert results[0].status is JobStatus.TRANSPORT_FAILED + assert protocol.calls == 1 + + +def test_coordinator_applies_timeout_to_the_entire_attempt(tmp_path): + coordinator = EnrollmentCoordinator( + source=Source(), + protocol=SlowDeviceFlowProtocol(), + executor=SlowExecutor(), + sink=Sink(), + ledger_path=tmp_path / "ledger.db", + ledger_salt=b"salt", + concurrency=1, + timeout=0.06, + retry_attempts=3, + ) + + results = asyncio.run(coordinator.run(target=1)) + + assert results[0].status is JobStatus.TIMEOUT + + +@pytest.mark.parametrize("source_kind", ["file", "sqlite"]) +def test_coordinator_consumes_real_source_adapters(tmp_path, source_kind): + if source_kind == "file": + source_path = tmp_path / "sessions.tsv" + source_path.write_text("file-source\tfile-token\n", encoding="utf-8") + source_path.chmod(0o600) + source = FileSourceAdapter(source_path) + else: + source_path = tmp_path / "accounts.db" + connection = sqlite3.connect(source_path) + connection.execute( + "CREATE TABLE accounts (token TEXT, status TEXT, deleted_at TEXT, updated_at INTEGER)" + ) + connection.execute("INSERT INTO accounts VALUES (?, ?, ?, ?)", ("db-token", "active", None, 1)) + connection.commit() + connection.close() + source = SQLiteSourceAdapter(source_path, b"source-salt") + + coordinator = EnrollmentCoordinator( + source=source, + protocol=Protocol(), + executor=Executor(), + sink=Sink(), + ledger_path=tmp_path / f"{source_kind}-ledger.db", + ledger_salt=b"ledger-salt", + concurrency=1, + timeout=5, + ) + results = asyncio.run(coordinator.run(target=1)) + assert len(results) == 1 + assert results[0].status is JobStatus.IMPORTED diff --git a/tests/test_xai_enroller_executors.py b/tests/test_xai_enroller_executors.py new file mode 100644 index 0000000..31faffc --- /dev/null +++ b/tests/test_xai_enroller_executors.py @@ -0,0 +1,411 @@ +import asyncio +from contextlib import suppress + +import httpx +import pytest + +from xai_enroller.executors import HTTPProbeExecutor, PlaywrightExecutor +from xai_enroller.models import AuthorizationStatus, DeviceFlow, SourceRecord + + +def test_http_probe_classifies_403_challenge_without_submission(): + calls = [] + + def handler(request): + calls.append(request) + return httpx.Response(403, text="challenge-platform") + + executor = HTTPProbeExecutor( + httpx.AsyncClient(transport=httpx.MockTransport(handler), follow_redirects=False) + ) + source = SourceRecord("source", "sso-token") + flow = DeviceFlow("device", "ABCD", "https://accounts.x.ai/oauth2/device?user_code=ABCD", 60, 1) + result = asyncio.run(executor.confirm(source, flow)) + assert result.status is AuthorizationStatus.NEEDS_BROWSER + assert len(calls) == 1 + assert calls[0].method == "GET" + assert "sso-token" not in repr(result) + + +def test_http_probe_does_not_follow_redirects_or_submit_forms(): + def handler(request): + return httpx.Response(302, headers={"location": "https://evil.example/"}) + + executor = HTTPProbeExecutor( + httpx.AsyncClient(transport=httpx.MockTransport(handler), follow_redirects=False) + ) + result = asyncio.run( + executor.confirm( + SourceRecord("source", "token"), + DeviceFlow("device", "ABCD", "https://accounts.x.ai/oauth2/device", 60, 1), + ) + ) + assert result.status is AuthorizationStatus.NEEDS_INTERACTION + + +class FakeButton: + def __init__(self, count): + self._count = count + + async def count(self): + return self._count + + @property + def first(self): + return self + + async def click(self): + return None + + +class FakePage: + def __init__(self): + self.closed = False + + async def goto(self, url, wait_until): + self.url = url + + def locator(self, selector): + return self + + async def inner_text(self): + return "Authorize access" + + def get_by_role(self, role, name): + return FakeButton(1 if name == "Authorize" else 0) + + async def close(self): + self.closed = True + + +class FakeContext: + def __init__(self): + self.cookies = [] + self.page = FakePage() + self.closed = False + + async def add_cookies(self, cookies): + self.cookies.extend(cookies) + + async def new_page(self): + return self.page + + async def close(self): + self.closed = True + + +class FakeBrowser: + def __init__(self): + self.launches = 0 + self.contexts = [] + self.closed = False + + async def new_context(self): + context = FakeContext() + self.contexts.append(context) + return context + + async def close(self): + self.closed = True + + +class FakeChromium: + def __init__(self, browser): + self.browser = browser + self.launch_options = None + + async def launch(self, **options): + self.launch_options = options + self.browser.launches += 1 + return self.browser + + +class FakePlaywright: + def __init__(self, browser): + self.chromium = FakeChromium(browser) + self.stopped = False + + async def stop(self): + self.stopped = True + + +class FakePlaywrightFactory: + def __init__(self): + self.browser = FakeBrowser() + self.playwright = FakePlaywright(self.browser) + + def __call__(self): + factory = self + + class Runner: + async def start(self): + return factory.playwright + + return Runner() + + +class FailingPageContext(FakeContext): + async def new_page(self): + raise RuntimeError("page creation failed") + + +class FailingPageBrowser(FakeBrowser): + async def new_context(self): + context = FailingPageContext() + self.contexts.append(context) + return context + + +def test_playwright_executor_isolates_context_and_scopes_exact_cookie(): + factory = FakePlaywrightFactory() + executor = PlaywrightExecutor(playwright_factory=factory) + source = SourceRecord("source", "secret-token") + flow = DeviceFlow("device", "ABCD", "https://accounts.x.ai/oauth2/device", 60, 1) + result = asyncio.run(executor.confirm(source, flow)) + asyncio.run(executor.close()) + assert result.status is AuthorizationStatus.AUTHORIZED + assert factory.browser.launches == 1 + context = factory.browser.contexts[0] + assert context.cookies == [{ + "name": "sso", + "value": "secret-token", + "domain": "accounts.x.ai", + "path": "/", + "secure": True, + "httpOnly": True, + }] + assert context.page.closed + assert context.closed + assert factory.browser.closed + + +def test_playwright_executor_closes_context_when_page_creation_fails(): + factory = FakePlaywrightFactory() + factory.browser = FailingPageBrowser() + factory.playwright = FakePlaywright(factory.browser) + executor = PlaywrightExecutor(playwright_factory=factory) + + result = asyncio.run( + executor.confirm( + SourceRecord("source", "secret-token"), + DeviceFlow("device", "ABCD", "https://accounts.x.ai/oauth2/device", 60, 1), + ) + ) + asyncio.run(executor.close()) + + assert result.status is AuthorizationStatus.NEEDS_INTERACTION + assert factory.browser.contexts[0].closed + + +def test_playwright_executor_launches_configured_fingerprint_browser(): + factory = FakePlaywrightFactory() + executor = PlaywrightExecutor( + playwright_factory=factory, + executable_path="/opt/cloakbrowser/chrome", + ) + + asyncio.run(executor.start()) + asyncio.run(executor.close()) + + assert factory.playwright.chromium.launch_options == { + "headless": True, + "executable_path": "/opt/cloakbrowser/chrome", + } + + +def test_playwright_executor_uses_fingerprint_browser_path_from_environment(monkeypatch): + monkeypatch.setenv("XAI_ENROLLER_BROWSER_EXECUTABLE", "/opt/cloakbrowser/chrome") + factory = FakePlaywrightFactory() + executor = PlaywrightExecutor(playwright_factory=factory) + + asyncio.run(executor.start()) + asyncio.run(executor.close()) + + assert factory.playwright.chromium.launch_options == { + "headless": True, + "executable_path": "/opt/cloakbrowser/chrome", + } + + +def test_playwright_executor_finds_macos_cloakbrowser_binary(monkeypatch): + macos_binary = ( + "/Users/test/.cloakbrowser/chromium-145/Chromium.app/Contents/MacOS/Chromium" + ) + monkeypatch.delenv("XAI_ENROLLER_BROWSER_EXECUTABLE", raising=False) + monkeypatch.setattr( + "xai_enroller.executors.glob.glob", + lambda pattern: [macos_binary] if "Chromium.app" in pattern else [], + ) + + assert PlaywrightExecutor._find_executable_path() == macos_binary + + +def test_playwright_attempt_cleanup_defers_repeated_cancellation(): + class Page: + def __init__(self): + self.entered = asyncio.Event() + self.release = asyncio.Event() + self.closed = False + + async def close(self): + self.entered.set() + await self.release.wait() + self.closed = True + + class Context: + def __init__(self): + self.closed = False + + async def close(self): + self.closed = True + + async def scenario(): + page = Page() + context = Context() + task = asyncio.create_task( + PlaywrightExecutor._close_attempt_resources(page, context) + ) + await page.entered.wait() + task.cancel() + task.cancel() + await asyncio.sleep(0) + page.release.set() + with pytest.raises(asyncio.CancelledError): + await task + assert page.closed + assert context.closed + + asyncio.run(scenario()) + + +def test_playwright_attempt_cleanup_has_a_true_wall_clock_deadline(monkeypatch): + class Page: + def __init__(self): + self.entered = asyncio.Event() + self.release = asyncio.Event() + + async def close(self): + self.entered.set() + while not self.release.is_set(): + try: + await self.release.wait() + except asyncio.CancelledError: + # Playwright transports can remain stuck while cancellation + # propagates. A hard deadline must not await that forever. + continue + + async def scenario(): + monkeypatch.setattr(PlaywrightExecutor, "CLOSE_TIMEOUT_SECONDS", 0.02) + page = Page() + task = asyncio.create_task( + PlaywrightExecutor._close_attempt_resources(page, None) + ) + await page.entered.wait() + done, _pending = await asyncio.wait({task}, timeout=0.10) + completed_within_deadline = task in done + + page.release.set() + if not task.done(): + task.cancel() + with suppress(asyncio.CancelledError): + await task + assert completed_within_deadline + + asyncio.run(scenario()) + + +def test_playwright_confirmation_has_hard_deadline_and_restarts_browser(monkeypatch): + class StubbornGotoPage(FakePage): + def __init__(self): + super().__init__() + self.entered = asyncio.Event() + self.release = asyncio.Event() + + async def goto(self, url, wait_until): + self.url = url + self.entered.set() + while not self.release.is_set(): + try: + await self.release.wait() + except asyncio.CancelledError: + # Model a wedged Playwright transport that does not finish + # propagating task cancellation. + continue + + class StubbornGotoContext(FakeContext): + def __init__(self): + super().__init__() + self.page = StubbornGotoPage() + + class StubbornGotoBrowser(FakeBrowser): + def __init__(self): + super().__init__() + self.close_entered = asyncio.Event() + self.close_release = asyncio.Event() + + async def new_context(self): + context = StubbornGotoContext() + self.contexts.append(context) + return context + + async def close(self): + self.close_entered.set() + while not self.close_release.is_set(): + try: + await self.close_release.wait() + except asyncio.CancelledError: + continue + self.closed = True + + class RotatingPlaywrightFactory: + def __init__(self): + self.browsers = [StubbornGotoBrowser(), FakeBrowser()] + self.playwrights = [FakePlaywright(browser) for browser in self.browsers] + self.starts = 0 + + def __call__(self): + factory = self + + class Runner: + async def start(self): + playwright = factory.playwrights[factory.starts] + factory.starts += 1 + return playwright + + return Runner() + + async def scenario(): + monkeypatch.setattr(PlaywrightExecutor, "ATTEMPT_TIMEOUT_SECONDS", 0.02) + monkeypatch.setattr(PlaywrightExecutor, "CLOSE_TIMEOUT_SECONDS", 0.02) + factory = RotatingPlaywrightFactory() + executor = PlaywrightExecutor(playwright_factory=factory) + source = SourceRecord("source", "secret-token") + flow = DeviceFlow( + "device", "ABCD", "https://accounts.x.ai/oauth2/device", 60, 1 + ) + + first = asyncio.create_task(executor.confirm(source, flow)) + while not factory.browsers[0].contexts: + await asyncio.sleep(0) + stuck_page = factory.browsers[0].contexts[0].page + await stuck_page.entered.wait() + done, _pending = await asyncio.wait({first}, timeout=0.10) + completed_within_deadline = first in done + + stuck_page.release.set() + factory.browsers[0].close_release.set() + first_result = await first + await asyncio.sleep(0) + second_result = await executor.confirm(source, flow) + await executor.close() + + assert completed_within_deadline + assert first_result.status is AuthorizationStatus.NEEDS_INTERACTION + assert first_result.reason_code == "confirmation_timeout" + assert second_result.status is AuthorizationStatus.AUTHORIZED + assert factory.starts == 2 + assert factory.browsers[0].close_entered.is_set() + assert factory.browsers[0].closed + assert factory.playwrights[0].stopped + + asyncio.run(scenario()) diff --git a/tests/test_xai_enroller_integration.py b/tests/test_xai_enroller_integration.py new file mode 100644 index 0000000..634718a --- /dev/null +++ b/tests/test_xai_enroller_integration.py @@ -0,0 +1,188 @@ +import asyncio +import json + +import httpx + +from xai_enroller.coordinator import EnrollmentCoordinator +from xai_enroller.executors import PlaywrightExecutor +from xai_enroller.models import SourceRecord +from xai_enroller.protocol import XAIProfile, XAIProtocol +from xai_enroller.sinks import CPAAuthFileSink + + +class SingleSource: + async def records(self): + yield SourceRecord("source-1", "sso-secret") + + +class FakeButton: + async def count(self): + return 1 + + @property + def first(self): + return self + + async def click(self): + return None + + +class FakePage: + async def goto(self, url, wait_until): + self.url = url + + def locator(self, selector): + return self + + async def inner_text(self): + return "Authorize access" + + def get_by_role(self, role, name): + return FakeButton() if name == "Authorize" else type( + "NoButton", + (), + {"count": staticmethod(lambda: asyncio.sleep(0, result=0))}, + )() + + async def close(self): + return None + + +class FakeContext: + def __init__(self, browser): + self.browser = browser + self.cookies = [] + + async def add_cookies(self, cookies): + self.cookies.extend(cookies) + + async def new_page(self): + return FakePage() + + async def close(self): + self.browser.live_contexts -= 1 + + +class FakeBrowser: + def __init__(self): + self.launches = 0 + self.live_contexts = 0 + self.max_live_contexts = 0 + + async def new_context(self): + self.live_contexts += 1 + self.max_live_contexts = max(self.max_live_contexts, self.live_contexts) + return FakeContext(self) + + async def close(self): + return None + + +class FakeChromium: + def __init__(self, browser): + self.browser = browser + + async def launch(self, **options): + self.browser.launches += 1 + return self.browser + + +class FakePlaywright: + def __init__(self, browser): + self.chromium = FakeChromium(browser) + + async def stop(self): + return None + + +class FakePlaywrightFactory: + def __init__(self): + self.browser = FakeBrowser() + self.playwright = FakePlaywright(self.browser) + + def __call__(self): + factory = self + + class Runner: + async def start(self): + return factory.playwright + + return Runner() + + +def test_adapter_boundary_hands_xai_credential_to_cpa_without_ledger_secrets(tmp_path): + def xai_handler(request): + if request.url.path == "/.well-known/openid-configuration": + return httpx.Response( + 200, + json={ + "device_authorization_endpoint": "https://auth.x.ai/oauth2/device/code", + "token_endpoint": "https://auth.x.ai/oauth2/token", + }, + ) + if request.url.path == "/oauth2/device/code": + return httpx.Response( + 200, + json={ + "device_code": "device-secret", + "user_code": "ABCD", + "verification_uri": "https://accounts.x.ai/oauth2/device", + "verification_uri_complete": ( + "https://accounts.x.ai/oauth2/device?user_code=ABCD" + ), + "expires_in": 60, + "interval": 0, + }, + ) + if request.url.path == "/oauth2/token": + return httpx.Response( + 200, + json={ + "access_token": "access-secret", + "refresh_token": "refresh-secret", + "id_token": "id-secret", + "token_type": "Bearer", + "expires_in": 3600, + "sub": "subject-1", + }, + ) + raise AssertionError(f"unexpected xAI request: {request.url}") + + cpa_requests = [] + + def cpa_handler(request): + cpa_requests.append(request) + return httpx.Response(201) + + xai_client = httpx.AsyncClient(transport=httpx.MockTransport(xai_handler)) + cpa_client = httpx.AsyncClient(transport=httpx.MockTransport(cpa_handler)) + factory = FakePlaywrightFactory() + coordinator = EnrollmentCoordinator( + source=SingleSource(), + protocol=XAIProtocol(xai_client, XAIProfile.default()), + executor=PlaywrightExecutor(playwright_factory=factory), + sink=CPAAuthFileSink( + "https://cpa.example", + "management-secret", + cpa_client, + name_secret=b"name-secret", + ), + ledger_path=tmp_path / "ledger.db", + ledger_salt=b"ledger-salt", + ) + + try: + results = asyncio.run(coordinator.run(target=1)) + finally: + asyncio.run(xai_client.aclose()) + asyncio.run(cpa_client.aclose()) + + assert results[0].status.value == "imported" + assert factory.browser.launches == 1 + assert factory.browser.max_live_contexts == 1 + assert cpa_requests[0].url.path == "/v0/management/auth-files" + assert cpa_requests[0].url.params["name"].startswith("xai-") + assert json.loads(cpa_requests[0].content)["refresh_token"] == "refresh-secret" + ledger_bytes = (tmp_path / "ledger.db").read_bytes() + for secret in ("sso-secret", "device-secret", "access-secret", "refresh-secret", "id-secret"): + assert secret.encode() not in ledger_bytes diff --git a/tests/test_xai_enroller_ledger.py b/tests/test_xai_enroller_ledger.py new file mode 100644 index 0000000..4ad80d0 --- /dev/null +++ b/tests/test_xai_enroller_ledger.py @@ -0,0 +1,163 @@ +import sqlite3 + +import pytest + +from xai_enroller.ledger import Ledger +from xai_enroller.models import JobStatus + + +def test_ledger_persists_only_redacted_terminal_fields(tmp_path): + path = tmp_path / "ledger.db" + ledger = Ledger(path, b"salt") + job_id = ledger.start("source", attempt=1) + ledger.finish(job_id, JobStatus.SINK_FAILED, "sink_failed", "receipt") + raw = path.read_bytes() + for secret in [ + "sso-token", + "device-code", + "https://accounts.x.ai", + "access-token", + "refresh-token", + "id-token", + "person@example.com", + ]: + assert secret.encode() not in raw + row = ledger.get(job_id) + assert row["status"] == "sink_failed" + assert row["reason_code"] == "sink_failed" + assert row["sink_receipt_fingerprint"] == "receipt" + assert "source" not in repr(row["source_fingerprint"]) + + +def test_ledger_recovers_pending_jobs(tmp_path): + ledger = Ledger(tmp_path / "ledger.db", b"salt") + job_id = ledger.start("source", attempt=1) + ledger.recover_pending() + assert ledger.get(job_id)["status"] == JobStatus.CANCELLED.value + + +def test_ledger_backfills_imported_receipts_into_available_inventory(tmp_path): + path = tmp_path / "ledger.db" + with sqlite3.connect(path) as connection: + connection.execute( + """ + CREATE TABLE jobs ( + job_id INTEGER PRIMARY KEY, + source_fingerprint TEXT NOT NULL, + attempt_number INTEGER NOT NULL, + status TEXT NOT NULL, + started_at TEXT NOT NULL, + finished_at TEXT, + reason_code TEXT, + sink_receipt_fingerprint TEXT + ) + """ + ) + connection.executemany( + "INSERT INTO jobs(source_fingerprint, attempt_number, status, started_at, " + "finished_at, reason_code, sink_receipt_fingerprint) VALUES (?, 1, ?, ?, ?, ?, ?)", + [ + ("source-old", "imported", "2026-01-01", "2026-01-02", "imported", "receipt-old"), + ("source-new", "imported", "2026-01-03", "2026-01-04", "imported", "receipt-new"), + ("source-failed", "sink_failed", "2026-01-05", "2026-01-06", "sink_failed", "receipt-failed"), + ], + ) + + ledger = Ledger(path, b"salt") + + assert ledger.inventory_counts() == { + "available": 2, + "claiming": 0, + "claimed": 0, + } + assert ledger.claim_available(2, "batch-backfill") == [ + "receipt-new", + "receipt-old", + ] + + +def test_imported_finish_and_inventory_insert_are_atomic(tmp_path): + ledger = Ledger(tmp_path / "ledger.db", b"salt") + job_id = ledger.start("source", attempt=1) + with sqlite3.connect(ledger.path) as connection: + connection.execute( + """ + CREATE TRIGGER reject_inventory_insert + BEFORE INSERT ON credential_inventory + BEGIN + SELECT RAISE(ABORT, 'inventory insert rejected'); + END + """ + ) + + with pytest.raises(sqlite3.IntegrityError, match="inventory insert rejected"): + ledger.finish(job_id, JobStatus.IMPORTED, "imported", "receipt") + + assert ledger.get(job_id)["status"] == "pending" + assert ledger.inventory_counts()["available"] == 0 + + +def test_claim_available_uses_latest_finish_and_tracks_batch(tmp_path): + ledger = Ledger(tmp_path / "ledger.db", b"salt") + jobs = [] + for source, receipt in [ + ("source-a", "receipt-a"), + ("source-b", "receipt-b"), + ("source-c", "receipt-c"), + ]: + job_id = ledger.start(source) + ledger.finish(job_id, JobStatus.IMPORTED, "imported", receipt) + jobs.append(job_id) + with sqlite3.connect(ledger.path) as connection: + connection.executemany( + "UPDATE jobs SET finished_at=? WHERE job_id=?", + [ + ("2026-01-01T00:00:00+00:00", jobs[0]), + ("2026-01-03T00:00:00+00:00", jobs[1]), + ("2026-01-02T00:00:00+00:00", jobs[2]), + ], + ) + + assert ledger.claim_available(2, "batch-1") == ["receipt-b", "receipt-c"] + assert ledger.claim_available(2, "batch-2") == ["receipt-a"] + assert ledger.pending_claims("batch-1") == [ + { + "sink_receipt_fingerprint": "receipt-b", + "batch_id": "batch-1", + "note": "", + }, + { + "sink_receipt_fingerprint": "receipt-c", + "batch_id": "batch-1", + "note": "", + }, + ] + assert ledger.inventory_counts() == { + "available": 0, + "claiming": 3, + "claimed": 0, + } + + +def test_claim_completion_and_recovery_are_batch_scoped(tmp_path): + ledger = Ledger(tmp_path / "ledger.db", b"salt") + for source, receipt in [ + ("source-a", "receipt-a"), + ("source-b", "receipt-b"), + ("source-c", "receipt-c"), + ]: + job_id = ledger.start(source) + ledger.finish(job_id, JobStatus.IMPORTED, "imported", receipt) + + assert ledger.claim_available(2, "batch-complete") + assert ledger.claim_available(1, "batch-recover") + assert ledger.mark_claimed("batch-complete", "delivered") == 2 + assert ledger.recover_claiming("batch-recover", note="worker_restart") == 1 + + assert ledger.pending_claims() == [] + assert ledger.inventory_counts() == { + "available": 1, + "claiming": 0, + "claimed": 2, + } + assert ledger.claim_available(1, "batch-retry") == ["receipt-a"] diff --git a/tests/test_xai_enroller_protocol.py b/tests/test_xai_enroller_protocol.py new file mode 100644 index 0000000..ec5fa6e --- /dev/null +++ b/tests/test_xai_enroller_protocol.py @@ -0,0 +1,245 @@ +import base64 +import json + +import httpx +import pytest + +from xai_enroller.models import AuthorizationStatus +from xai_enroller.protocol import XAIProtocol, XAIProfile + + +def _unsigned_jwt(payload): + encoded = base64.urlsafe_b64encode( + json.dumps(payload, separators=(",", ":")).encode() + ).decode().rstrip("=") + return f"header.{encoded}.signature" + + +def test_default_profile_matches_observed_grok_cli_device_contract(): + profile = XAIProfile.default() + + assert profile.client_id == "b1a00492-073a-47ea-816f-4c329264a828" + assert profile.scope == "openid profile email offline_access grok-cli:access api:access" + assert profile.version == "grok-cli-device-v1" + + +def test_device_flow_uses_discovered_endpoints_and_profile_scope(): + requests = [] + + def handler(request): + requests.append(request) + if request.url.path == "/.well-known/openid-configuration": + return httpx.Response( + 200, + json={ + "device_authorization_endpoint": "https://auth.x.ai/oauth2/device/code", + "token_endpoint": "https://auth.x.ai/oauth2/token", + }, + ) + return httpx.Response( + 200, + json={ + "device_code": "device", + "user_code": "ABCD", + "verification_uri": "https://accounts.x.ai/oauth2/device", + "expires_in": 60, + "interval": 1, + }, + ) + + profile = XAIProfile(client_id="client", scope="openid offline_access") + protocol = XAIProtocol(httpx.AsyncClient(transport=httpx.MockTransport(handler)), profile) + + async def run(): + flow = await protocol.start_device_flow() + return flow + + flow = __import__("asyncio").run(run()) + assert flow.verification_url == "https://accounts.x.ai/oauth2/device?user_code=ABCD" + assert dict(__import__("urllib.parse", fromlist=["parse_qsl"]).parse_qsl(requests[1].content.decode())) == { + "client_id": "client", + "scope": "openid offline_access", + } + + +def test_device_flow_uses_configured_poll_interval_when_issuer_omits_it(): + def handler(request): + if request.url.path == "/.well-known/openid-configuration": + return httpx.Response( + 200, + json={ + "device_authorization_endpoint": "https://auth.x.ai/oauth2/device/code", + "token_endpoint": "https://auth.x.ai/oauth2/token", + }, + ) + return httpx.Response( + 200, + json={ + "device_code": "device", + "user_code": "ABCD", + "verification_uri": "https://accounts.x.ai/oauth2/device", + "expires_in": 60, + }, + ) + + protocol = XAIProtocol( + httpx.AsyncClient(transport=httpx.MockTransport(handler)), + XAIProfile(client_id="client", scope="openid"), + default_poll_interval=7, + ) + + flow = __import__("asyncio").run(protocol.start_device_flow()) + + assert flow.interval == 7 + + +def test_device_flow_preserves_complete_verification_url(): + def handler(request): + if request.url.path == "/.well-known/openid-configuration": + return httpx.Response( + 200, + json={ + "device_authorization_endpoint": "https://auth.x.ai/oauth2/device/code", + "token_endpoint": "https://auth.x.ai/oauth2/token", + }, + ) + return httpx.Response( + 200, + json={ + "device_code": "device", + "user_code": "ABCD", + "verification_uri": "https://accounts.x.ai/oauth2/device", + "verification_uri_complete": "https://accounts.x.ai/oauth2/device?user_code=ABCD&state=issuer", + "expires_in": 60, + }, + ) + + protocol = XAIProtocol( + httpx.AsyncClient(transport=httpx.MockTransport(handler)), + XAIProfile(client_id="client", scope="openid"), + ) + + flow = __import__("asyncio").run(protocol.start_device_flow()) + + assert flow.verification_url == ( + "https://accounts.x.ai/oauth2/device?user_code=ABCD&state=issuer" + ) + + +def test_xai_url_allowlist_rejects_missing_hostname(): + assert XAIProtocol._allowed_url("https:///oauth2/device") is False + + +@pytest.mark.parametrize( + "discovery", + [ + { + "device_authorization_endpoint": "http://auth.x.ai/device", + "token_endpoint": "https://auth.x.ai/token", + }, + { + "device_authorization_endpoint": "https://evil.example/device", + "token_endpoint": "https://auth.x.ai/token", + }, + ], +) +def test_discovery_rejects_non_xai_or_non_https(discovery): + def handler(request): + return httpx.Response(200, json=discovery) + + protocol = XAIProtocol( + httpx.AsyncClient(transport=httpx.MockTransport(handler)), + XAIProfile(client_id="client", scope="openid"), + ) + with pytest.raises(ValueError, match="allowlist"): + __import__("asyncio").run(protocol.discover()) + + +def test_token_polling_maps_pending_slow_down_and_denial(): + responses = iter( + [ + httpx.Response(400, json={"error": "authorization_pending"}), + httpx.Response(400, json={"error": "slow_down"}), + httpx.Response(400, json={"error": "access_denied"}), + ] + ) + + def handler(request): + return next(responses) + + protocol = XAIProtocol( + httpx.AsyncClient(transport=httpx.MockTransport(handler)), + XAIProfile(client_id="client", scope="openid"), + sleep=lambda _: __import__("asyncio").sleep(0), + ) + with pytest.raises(RuntimeError, match="oauth_denied"): + __import__("asyncio").run( + protocol.poll_token( + endpoint="https://auth.x.ai/oauth2/token", + flow=type("Flow", (), {"device_code": "device", "interval": 0})(), + timeout=1, + ) + ) + + +def test_token_response_requires_refresh_token_and_never_exposes_body(): + def handler(request): + return httpx.Response(200, json={"access_token": "access", "token_type": "Bearer"}) + + protocol = XAIProtocol( + httpx.AsyncClient(transport=httpx.MockTransport(handler)), + XAIProfile(client_id="client", scope="openid"), + ) + with pytest.raises(RuntimeError) as error: + __import__("asyncio").run( + protocol.poll_token( + endpoint="https://auth.x.ai/oauth2/token", + flow=type("Flow", (), {"device_code": "device", "interval": 0})(), + timeout=1, + ) + ) + assert "access" not in str(error.value) + + +def test_credential_subject_prefers_id_token_subject(): + credential = XAIProtocol._credential( + { + "access_token": _unsigned_jwt({"principal_id": "access-subject"}), + "refresh_token": "refresh", + "id_token": _unsigned_jwt( + {"sub": "stable-subject", "email": "ignored@example.test"} + ), + "sub": "response-subject", + }, + "https://auth.x.ai/oauth2/token", + ) + + assert credential.subject == "stable-subject" + + +def test_credential_subject_falls_back_to_access_token_principal_id(): + credential = XAIProtocol._credential( + { + "access_token": _unsigned_jwt({"principal_id": "stable-principal"}), + "refresh_token": "refresh", + "id_token": "malformed", + "sub": "response-subject", + }, + "https://auth.x.ai/oauth2/token", + ) + + assert credential.subject == "stable-principal" + + +def test_credential_subject_ignores_email_and_falls_back_to_response_subject(): + credential = XAIProtocol._credential( + { + "access_token": "not-a-jwt", + "refresh_token": "refresh", + "id_token": _unsigned_jwt({"email": "ignored@example.test"}), + "sub": "response-subject", + }, + "https://auth.x.ai/oauth2/token", + ) + + assert credential.subject == "response-subject" diff --git a/tests/test_xai_enroller_sinks.py b/tests/test_xai_enroller_sinks.py new file mode 100644 index 0000000..65e61b2 --- /dev/null +++ b/tests/test_xai_enroller_sinks.py @@ -0,0 +1,79 @@ +import asyncio +import json + +import httpx + +from xai_enroller.models import OAuthCredential +from xai_enroller.sinks import CPAAuthFileSink, SinkError + + +def credential(): + return OAuthCredential( + access_token="access", + refresh_token="refresh", + id_token="id-token", + token_type="Bearer", + expires_in=3600, + expires_at="2026-07-11T00:00:00Z", + last_refresh="2026-07-11T00:00:00Z", + subject="subject", + token_endpoint="https://auth.x.ai/oauth2/token", + ) + + +def test_cpa_sink_builds_pinned_document_and_safe_name(): + requests = [] + + def handler(request): + requests.append(request) + return httpx.Response(201) + + sink = CPAAuthFileSink( + "https://cpa.example", + "management-secret", + httpx.AsyncClient(transport=httpx.MockTransport(handler)), + name_secret=b"name-secret", + ) + receipt = asyncio.run(sink.store(credential())) + assert requests[0].url.path == "/v0/management/auth-files" + assert requests[0].headers["authorization"] == "Bearer management-secret" + assert requests[0].url.params["name"].startswith("xai-") + assert requests[0].url.params["name"].endswith(".json") + assert json.loads(requests[0].content) == { + "type": "xai", + "access_token": "access", + "refresh_token": "refresh", + "id_token": "id-token", + "token_type": "Bearer", + "expires_in": 3600, + "expired": "2026-07-11T00:00:00Z", + "last_refresh": "2026-07-11T00:00:00Z", + "sub": "subject", + "base_url": "https://api.x.ai/v1", + "token_endpoint": "https://auth.x.ai/oauth2/token", + "auth_kind": "oauth", + } + assert "access" not in repr(receipt) + + +def test_cpa_sink_maps_non_2xx_without_retry_or_secret_leak(): + count = 0 + + def handler(request): + nonlocal count + count += 1 + return httpx.Response(500, text="access") + + sink = CPAAuthFileSink( + "https://cpa.example", + "management-secret", + httpx.AsyncClient(transport=httpx.MockTransport(handler)), + name_secret=b"name-secret", + ) + try: + asyncio.run(sink.store(credential())) + except SinkError as error: + assert "access" not in str(error) + else: + raise AssertionError("expected SinkError") + assert count == 1 diff --git a/tests/test_xai_remote_snapshot.py b/tests/test_xai_remote_snapshot.py new file mode 100644 index 0000000..c876f1d --- /dev/null +++ b/tests/test_xai_remote_snapshot.py @@ -0,0 +1,240 @@ +import asyncio +import json +import os + +from xai_enroller.models import SourceRecord +from xai_enroller.remote_stream import ( + DiskSnapshotSource, + LocalSnapshotSynchronizer, + SSHSnapshotSynchronizer, +) + + +def _document(source_id): + return json.dumps( + { + "email": source_id, + "cookies": [ + { + "name": "sso", + "value": "opaque", + "domain": "accounts.x.ai", + "path": "/", + } + ], + }, + separators=(",", ":"), + ).encode() + + +class FakeStream: + def __init__(self, lines): + self._lines = iter(lines) + + async def readline(self): + return next(self._lines, b"") + + async def read(self, _size=-1): + return b"" + + +class FakeProcess: + def __init__(self, lines, returncode=0): + self.stdout = FakeStream(lines) + self.stderr = FakeStream([]) + self.returncode = None + self._final_returncode = returncode + + async def wait(self): + self.returncode = self._final_returncode + return self.returncode + + def terminate(self): + self.returncode = -15 + + def kill(self): + self.returncode = -9 + + +def test_snapshot_sync_atomically_replaces_only_valid_complete_export(tmp_path): + async def scenario(): + destination = tmp_path / "source-snapshot.jsonl" + destination.write_bytes(_document("old") + b"\n") + + invalid = FakeProcess([b"not-json\n"]) + + async def invalid_factory(*_args, **_kwargs): + return invalid + + synchronizer = SSHSnapshotSynchronizer( + "host", + destination, + process_factory=invalid_factory, + fingerprint=lambda source_id: f"key:{source_id}", + ) + assert not await synchronizer.sync_once() + assert destination.read_bytes() == _document("old") + b"\n" + + rejected = FakeProcess([_document("rejected") + b"\n"], returncode=1) + + async def rejected_factory(*_args, **_kwargs): + return rejected + + synchronizer.process_factory = rejected_factory + assert not await synchronizer.sync_once() + assert destination.read_bytes() == _document("old") + b"\n" + + empty = FakeProcess([]) + + async def empty_factory(*_args, **_kwargs): + return empty + + synchronizer.process_factory = empty_factory + assert not await synchronizer.sync_once() + assert destination.read_bytes() == _document("old") + b"\n" + + valid = FakeProcess([_document("new-a") + b"\n", _document("new-b") + b"\n"]) + + async def valid_factory(*_args, **_kwargs): + return valid + + synchronizer.process_factory = valid_factory + assert await synchronizer.sync_once() + assert synchronizer.snapshot_fingerprints == frozenset( + {"key:new-a", "key:new-b"} + ) + assert destination.read_bytes() == _document("new-a") + b"\n" + _document("new-b") + b"\n" + assert destination.stat().st_mode & 0o777 == 0o600 + assert not list(tmp_path.glob(".source-snapshot.jsonl.*.tmp")) + + synchronizer.process_factory = invalid_factory + assert not await synchronizer.sync_once() + assert synchronizer.snapshot_fingerprints == frozenset( + {"key:new-a", "key:new-b"} + ) + + asyncio.run(scenario()) + + +def test_local_snapshot_accepts_empty_input_and_anchors_exporter_code(tmp_path): + async def scenario(): + registration_root = tmp_path / "registration" + destination = tmp_path / "auth" / "source-snapshot.jsonl" + captured = [] + process = FakeProcess([]) + + async def factory(*args, **_kwargs): + captured.append(args) + return process + + synchronizer = LocalSnapshotSynchronizer( + registration_root, + destination, + process_factory=factory, + ) + + assert await synchronizer.sync_once() + assert destination.read_bytes() == b"" + assert destination.stat().st_mode & 0o777 == 0o600 + assert synchronizer.snapshot_fingerprints == frozenset() + assert captured[0][0] + assert captured[0][1] != str(registration_root / "scripts" / "export_registered_sessions.py") + assert captured[0][-2:] == ( + str(registration_root / "keys" / "auth-sessions.jsonl"), + str(registration_root / "keys" / "accounts.txt"), + ) + + asyncio.run(scenario()) + + +def test_local_snapshot_rejects_candidate_when_either_input_changes(tmp_path): + async def scenario(): + registration_root = tmp_path / "registration" + keys = registration_root / "keys" + keys.mkdir(parents=True) + sessions = keys / "auth-sessions.jsonl" + accounts = keys / "accounts.txt" + sessions.write_bytes(_document("old") + b"\n") + accounts.write_text("", encoding="utf-8") + destination = tmp_path / "auth" / "source-snapshot.jsonl" + destination.parent.mkdir() + destination.write_bytes(_document("stable") + b"\n") + + process = FakeProcess([_document("candidate") + b"\n"]) + + async def factory(*_args, **_kwargs): + accounts.write_text("new@example.test:p:sso\n", encoding="utf-8") + return process + + synchronizer = LocalSnapshotSynchronizer( + registration_root, + destination, + process_factory=factory, + ) + + assert not await synchronizer.sync_once() + assert destination.read_bytes() == _document("stable") + b"\n" + + asyncio.run(scenario()) + + +def test_snapshot_consumer_finishes_open_generation_before_replacement(tmp_path): + async def scenario(): + destination = tmp_path / "source-snapshot.jsonl" + destination.write_bytes(_document("old-a") + b"\n" + _document("old-b") + b"\n") + source = DiskSnapshotSource(destination, poll_seconds=0.01) + records = source.records() + first = await anext(records) + assert isinstance(first, SourceRecord) + assert first.source_id == "old-a" + + replacement = tmp_path / "replacement" + replacement.write_bytes(_document("new-a") + b"\n") + os.replace(replacement, destination) + + assert (await anext(records)).source_id == "old-b" + assert (await asyncio.wait_for(anext(records), 1)).source_id == "new-a" + await source.close() + await records.aclose() + + asyncio.run(scenario()) + + +def test_snapshot_source_reports_only_aggregate_new_records(tmp_path): + class Synchronizer: + def __init__(self): + self.snapshot_fingerprints = None + self.calls = 0 + self.source = None + + async def sync_once(self): + self.calls += 1 + if self.calls == 1: + self.snapshot_fingerprints = frozenset({"a", "b"}) + else: + self.snapshot_fingerprints = frozenset({"a", "b", "c"}) + self.source._closed = True + return True + + async def close(self): + return None + + async def scenario(): + synchronizer = Synchronizer() + events = [] + source = DiskSnapshotSource( + tmp_path / "source-snapshot.jsonl", + synchronizer=synchronizer, + sync_seconds=0.001, + event_callback=lambda kind, data: events.append((kind, data)), + ) + synchronizer.source = source + + await source._sync_loop() + + assert events == [ + ("source_connected", {}), + ("source_updated", {"new": 1, "total": 3}), + ] + + asyncio.run(scenario()) diff --git a/tools/GrokSession2CPAAndSub2API/README.md b/tools/GrokSession2CPAAndSub2API/README.md new file mode 100644 index 0000000..4309304 --- /dev/null +++ b/tools/GrokSession2CPAAndSub2API/README.md @@ -0,0 +1,100 @@ +# Grok Session → CPA / sub2api + +纯前端单页工具:把 grok-free-register 产出的凭证,转成 CPA、sub2api 导入包、batch 建号、或 sso-to-oauth 请求体。 + +对照思路:GPT session → CPA / sub2api 纯前端转换器(同类工具,无外部依赖)。 + +## 本地使用 + +直接打开: + +```text +tools/GrokSession2CPAAndSub2API/index.html +``` + +所有解析和转换都在浏览器本地完成,不上传 token。 + +## 支持输入 + +| 来源 | 说明 | +|------|------| +| `authenticated/xai-*.json` | 注册成功 OAuth(base_url 常是 api.x.ai) | +| `cpa_ready/xai-*.json` | watchdog 终态(带 headers/email) | +| `sub2api-data` / accounts 数组 | 已转换过的包,可再导成其它格式 | +| `auth-sessions.jsonl` | `email` + `cookies.sso` | +| `accounts.txt` | `email:password:sso` | +| `grok.txt` | 每行一个 SSO JWT | + +可粘贴 JSON / 数组 / jsonl,或拖入多个文件。 + +## 输出格式 + +### CPA + +对齐注册机 cpa_ready: + +- `type: "xai"` +- `access_token` / `refresh_token` / `id_token` +- `base_url` **强制** `https://cli-chat-proxy.grok.com/v1`(忽略 authenticated 里的 `api.x.ai`) +- `headers`:保留源 headers,否则写 grok-cli 默认头 +- `email` / `sub` / `expired` / `expires_in` + +单条输出对象,多条输出数组。 + +### sub2api 导入 + +`type: "sub2api-data"` 包,可走 sub2api 数据导入: + +```json +{ + "type": "sub2api-data", + "version": 1, + "exported_at": "...", + "proxies": [], + "accounts": [ ... ] +} +``` + +账号字段: + +- `platform: "grok"`, `type: "oauth"` +- `credentials.access_token` / `refresh_token` / `id_token` / `expires_at` / `email` / `sub` / `base_url` +- **有 `refresh_token`**:不写账号顶层 `expires_at` / `auto_pause_on_expired`(避免 access 6h 到期被 pause) +- **无 `refresh_token`**:写 access JWT exp + `auto_pause_on_expired: true` + +### batch 建号 + +`POST /api/v1/admin/accounts/batch` 体: + +```json +{ "accounts": [ /* 同上 sub2api 账号对象 */ ] } +``` + +### sso-to-oauth + +`POST /api/v1/admin/grok/sso-to-oauth` 体: + +```json +{ + "sso_tokens": ["eyJ...", "..."], + "_meta": [ /* 仅本地对照,API 可忽略 */ ] +} +``` + +纯 SSO 输入(grok.txt / accounts.txt / auth-sessions)**只能**走这个格式;CPA / sub2api / batch 需要 OAuth。 + +## 关键规则(别踩坑) + +1. free 号 API 走 `cli-chat-proxy.grok.com`,不是 `api.x.ai`。 +2. free OAuth access ≈ 6h(`expires_in=21600`)。有 refresh 时让 sub2api 自己刷;别把账号级 `expires_at` 钉成 access exp。 +3. 当前线上常见:refresh 被整批 revoke,access 在 exp 前仍可用。这种号导入后到期无法续命,只能重铸。 +4. 注册机 `grok.txt` / `accounts.txt` / `auth-sessions.jsonl` 只有 SSO,没有 OAuth;要进 sub2api 要么走 sso-to-oauth,要么先走 xai_enroller 出 `authenticated/` 再转。 + +## 与注册机衔接 + +``` +register → grok.txt / accounts.txt / auth-sessions.jsonl (SSO) + → xai_enroller → authenticated/xai-*.json (OAuth) + → acpa_watchdog → cpa_ready/xai-*.json + acc.md + → 本工具 → sub2api-data / batch / CPA / sso-to-oauth +``` diff --git a/tools/GrokSession2CPAAndSub2API/index.html b/tools/GrokSession2CPAAndSub2API/index.html new file mode 100644 index 0000000..0ba8038 --- /dev/null +++ b/tools/GrokSession2CPAAndSub2API/index.html @@ -0,0 +1,902 @@ + + + + + + Grok Session → CPA / sub2api + + + + +
+
+
+

Grok Free Register toolkit

+

Grok Session → CPA / sub2api

+

+ 纯前端转换:把注册机 xai-*.json / SSO / accounts / sessions + 转成 CPA、sub2api 导入包、batch 建号、或 sso-to-oauth 请求体。本地解析,不上传。 +

+
+
+ platform: grok + base: cli-chat-proxy + 有 refresh 不写账号 6h 死期 +
+
+ +
+
+ + + + +
+
+ + + + +
+
+ +
+
+
+
+

输入

+

粘贴 / 拖入 JSON、JSON 数组、jsonl、accounts 行、纯 SSO 行

+
+
+
+
+ 支持输入
+ 1. authenticated/xai-*.jsoncpa_ready/xai-*.json
+ 2. sub2api account / accounts 数组 / sub2api-data 包
+ 3. auth-sessions.jsonl(email + cookies.sso)
+ 4. accounts.txtemail:password:sso
+ 5. grok.txt:每行一个 SSO JWT +
+ +
+ 等待输入 + 未转换 +
+
+
+
+
+ +
+
+
+

输出

+

CPA:type=xai,对齐注册机 cpa_ready / new-api 风格

+
+
+
+ +
+ +
+ +
+
+
+
+ + + + diff --git a/tools/__init__.py b/tools/__init__.py new file mode 100644 index 0000000..fdaa1e9 --- /dev/null +++ b/tools/__init__.py @@ -0,0 +1 @@ +"""Developer-only analysis and stress tools.""" diff --git a/tools/runtime_log_analyzer.py b/tools/runtime_log_analyzer.py new file mode 100644 index 0000000..366306c --- /dev/null +++ b/tools/runtime_log_analyzer.py @@ -0,0 +1,419 @@ +""" +Parse real runtime monitor logs from both CSP and legacy state-machine runs. + +The analyzer is intentionally read-only. It does not tune parameters or modify +runtime state; it only turns production logs into comparable stage rates. +""" +from __future__ import annotations + +from dataclasses import dataclass +import json +import re +from typing import Iterable + + +_CSP_RE = re.compile( + r"^\[\*\] T:(?P\d+) Q:(?P\d+) phys:(?P\d+) " + r"(?:p_send:(?P\d+) )?t_slot:(?P\d+) q_slot:(?P\d+) q_pend:(?P\d+)" + r"(?P.*)rate:(?P[0-9.]+)/min #(?P\d+)" +) + +_STATE_RE = re.compile( + r"^\[\*\] slots:(?P\d+)/(?P\d+) act:(?P\d+) " + r"cpu:(?P[0-9.]+)% avg:(?P[0-9.]+)%/(?P[0-9.]+) " + r"mem:(?P\d+)M T:(?P\d+) Q:(?P\d+) " + r"sent:(?P\d+) got:(?P\d+)\((?P[0-9.]+)%\) " + r"rate:(?P[0-9.]+)/min #(?P\d+)" +) + +_SOLVER_TIMELINE_PREFIX = "[solver_timeline] " + + +@dataclass(frozen=True) +class MonitorRow: + kind: str + t: int + q: int + ok: int + rate: float + phys: int | None = None + p_send: int | None = None + t_slot: int | None = None + q_slot: int | None = None + q_pend: int | None = None + p_batch: float | None = None + t_prog: int | None = None + q_inflight: int | None = None + t_prod: int | None = None + q_sent: int | None = None + q_ret: int | None = None + q_adm: int | None = None + pair: int | None = None + fail: int | None = None + s_phys_wait: float | None = None + s_phys_hold: float | None = None + p_phys_wait: float | None = None + p_phys_hold: float | None = None + c_phys_wait: float | None = None + c_phys_hold: float | None = None + p_email_create: float | None = None + p_page_prepare: float | None = None + p_send_stage: float | None = None + c_page_acquire: float | None = None + c_verify: float | None = None + c_register: float | None = None + c_hot_hits: int | None = None + c_hot_misses: int | None = None + solver_goto: float | None = None + solver_inject: float | None = None + solver_initial: float | None = None + solver_click: float | None = None + solver_wait: float | None = None + solver_reuse: float | None = None + solver_visible: float | None = None + slots: int | None = None + max_slots: int | None = None + active: int | None = None + + @property + def elapsed_min(self) -> float | None: + if self.ok <= 0 or self.rate <= 0: + return None + return self.ok / self.rate + + +@dataclass(frozen=True) +class SolverTimeline: + events: list[dict] + + +def _int_field(rest: str, name: str) -> int | None: + match = re.search(rf"\b{name}:(\d+)", rest) + return int(match.group(1)) if match else None + + +def _float_field(rest: str, name: str) -> float | None: + match = re.search(rf"\b{name}:([0-9.]+)", rest) + return float(match.group(1)) if match else None + + +def _float_pair_field(rest: str, name: str) -> tuple[float | None, float | None]: + match = re.search(rf"\b{name}:([0-9.]+)/([0-9.]+)", rest) + if not match: + return None, None + return float(match.group(1)), float(match.group(2)) + + +def _float_triple_field(rest: str, name: str) -> tuple[float | None, float | None, float | None]: + match = re.search(rf"\b{name}:([0-9.]+)/([0-9.]+)/([0-9.]+)", rest) + if not match: + return None, None, None + return float(match.group(1)), float(match.group(2)), float(match.group(3)) + + +def _int_pair_field(rest: str, name: str) -> tuple[int | None, int | None]: + match = re.search(rf"\b{name}:(\d+)/(\d+)", rest) + if not match: + return None, None + return int(match.group(1)), int(match.group(2)) + + +def parse_monitor_lines(text_or_lines: str | Iterable[str]) -> list[MonitorRow]: + if isinstance(text_or_lines, str): + lines = text_or_lines.splitlines() + else: + lines = list(text_or_lines) + + rows: list[MonitorRow] = [] + for line in lines: + csp = _CSP_RE.match(line) + if csp: + rest = csp.group("rest") + s_phys_wait, s_phys_hold = _float_pair_field(rest, "s_phys") + p_phys_wait, p_phys_hold = _float_pair_field(rest, "p_phys") + c_phys_wait, c_phys_hold = _float_pair_field(rest, "c_phys") + p_email_create, p_page_prepare, p_send_stage = _float_triple_field(rest, "p_stage") + c_page_acquire, c_verify, c_register = _float_triple_field(rest, "c_stage") + c_hot_hits, c_hot_misses = _int_pair_field(rest, "c_hot") + rows.append( + MonitorRow( + kind="csp", + t=int(csp.group("t")), + q=int(csp.group("q")), + ok=int(csp.group("ok")), + rate=float(csp.group("rate")), + phys=int(csp.group("phys")), + p_send=int(csp.group("p_send")) if csp.group("p_send") is not None else None, + t_slot=int(csp.group("t_slot")), + q_slot=int(csp.group("q_slot")), + q_pend=int(csp.group("q_pend")), + p_batch=_float_field(rest, "p_batch"), + t_prog=_int_field(rest, "t_prog"), + q_inflight=_int_field(rest, "q_inflight"), + t_prod=_int_field(rest, "t_prod"), + q_sent=_int_field(rest, "q_sent"), + q_ret=_int_field(rest, "q_ret"), + q_adm=_int_field(rest, "q_adm"), + pair=_int_field(rest, "pair"), + fail=_int_field(rest, "fail"), + s_phys_wait=s_phys_wait, + s_phys_hold=s_phys_hold, + p_phys_wait=p_phys_wait, + p_phys_hold=p_phys_hold, + c_phys_wait=c_phys_wait, + c_phys_hold=c_phys_hold, + p_email_create=p_email_create, + p_page_prepare=p_page_prepare, + p_send_stage=p_send_stage, + c_page_acquire=c_page_acquire, + c_verify=c_verify, + c_register=c_register, + c_hot_hits=c_hot_hits, + c_hot_misses=c_hot_misses, + solver_goto=_float_field(rest, "solver_goto"), + solver_inject=_float_field(rest, "solver_inject"), + solver_initial=_float_field(rest, "solver_initial"), + solver_click=_float_field(rest, "solver_click"), + solver_wait=_float_field(rest, "solver_wait"), + solver_reuse=_float_field(rest, "solver_reuse"), + solver_visible=_float_field(rest, "solver_visible"), + ) + ) + continue + + state = _STATE_RE.match(line) + if state: + rows.append( + MonitorRow( + kind="state_machine", + t=int(state.group("t")), + q=int(state.group("q")), + ok=int(state.group("ok")), + rate=float(state.group("rate")), + q_sent=int(state.group("q_sent")), + q_ret=int(state.group("q_ret")), + slots=int(state.group("slots")), + max_slots=int(state.group("max_slots")), + active=int(state.group("active")), + ) + ) + return rows + + +def _rate_delta(rows: list[MonitorRow], attr: str) -> float | None: + candidates = [r for r in rows if r.elapsed_min is not None and getattr(r, attr) is not None] + if len(candidates) < 2: + return None + first = candidates[0] + last = candidates[-1] + dt = (last.elapsed_min or 0) - (first.elapsed_min or 0) + if dt <= 0: + return None + return (getattr(last, attr) - getattr(first, attr)) / dt + + +def _leader(values: dict[str, float | None]) -> str | None: + if any(value is None for value in values.values()): + return None + return max(values, key=lambda key: values[key] or 0) + + +def _avg(values: list[float]) -> float | None: + if not values: + return None + return round(sum(values) / len(values), 2) + + +def _ratio(numerator: int, denominator: int) -> float | None: + if denominator <= 0: + return None + return round(numerator / denominator, 3) + + +def summarize_monitor_rows(rows: list[MonitorRow], recent_count: int = 6) -> dict[str, float | int | str | None]: + if not rows: + return {"rows": 0} + + last = rows[-1] + recent = rows[-recent_count:] + summary: dict[str, float | int | str | None] = { + "rows": len(rows), + "kind": last.kind, + "last_ok": last.ok, + "last_cumulative_rate": last.rate, + "last_q_minus_t": last.q - last.t, + "last_phys": last.phys, + "last_p_send_sem": last.p_send, + "last_t_slot": last.t_slot, + "last_q_slot": last.q_slot, + "last_q_pend": last.q_pend, + "last_p_batch": last.p_batch, + "last_t_prog": last.t_prog, + "last_q_inflight": last.q_inflight, + "last_q_return_minus_t_prod": ( + last.q_ret - last.t_prod + if last.q_ret is not None and last.t_prod is not None + else None + ), + "last_slots": last.slots, + "last_t_prod": last.t_prod, + "last_q_ret": last.q_ret, + "last_q_adm": last.q_adm, + "last_pair": last.pair, + "last_fail": last.fail, + "last_s_phys_wait": last.s_phys_wait, + "last_s_phys_hold": last.s_phys_hold, + "last_p_phys_wait": last.p_phys_wait, + "last_p_phys_hold": last.p_phys_hold, + "last_c_phys_wait": last.c_phys_wait, + "last_c_phys_hold": last.c_phys_hold, + "last_p_email_create": last.p_email_create, + "last_p_page_prepare": last.p_page_prepare, + "last_p_send": last.p_send_stage, + "last_c_page_acquire": last.c_page_acquire, + "last_c_verify": last.c_verify, + "last_c_register": last.c_register, + "last_c_hot_hits": last.c_hot_hits, + "last_c_hot_misses": last.c_hot_misses, + "last_physical_hold_leader": _leader({ + "s": last.s_phys_hold, + "p": last.p_phys_hold, + "c": last.c_phys_hold, + }), + "last_physical_wait_leader": _leader({ + "s": last.s_phys_wait, + "p": last.p_phys_wait, + "c": last.c_phys_wait, + }), + "last_solver_goto": last.solver_goto, + "last_solver_inject": last.solver_inject, + "last_solver_initial": last.solver_initial, + "last_solver_click": last.solver_click, + "last_solver_wait": last.solver_wait, + "last_solver_reuse": last.solver_reuse, + "last_solver_visible": last.solver_visible, + "recent_ok_per_min": _rate_delta(recent, "ok"), + "recent_t_prod_per_min": _rate_delta(recent, "t_prod"), + "recent_q_ret_per_min": _rate_delta(recent, "q_ret"), + "recent_pair_per_min": _rate_delta(recent, "pair"), + } + return summary + + +def parse_solver_timelines(text_or_lines: str | Iterable[str]) -> list[SolverTimeline]: + if isinstance(text_or_lines, str): + lines = text_or_lines.splitlines() + else: + lines = list(text_or_lines) + + timelines: list[SolverTimeline] = [] + for line in lines: + if not line.startswith(_SOLVER_TIMELINE_PREFIX): + continue + payload = line[len(_SOLVER_TIMELINE_PREFIX):] + try: + events = json.loads(payload) + except json.JSONDecodeError: + continue + if isinstance(events, list): + timelines.append(SolverTimeline(events=[e for e in events if isinstance(e, dict)])) + return timelines + + +def _last_page_trace(events: list[dict]) -> dict | None: + for event in reversed(events): + page_trace = event.get("page_trace") + if isinstance(page_trace, dict): + return page_trace + return None + + +def summarize_solver_timelines(timelines: list[SolverTimeline]) -> dict[str, float | int | None]: + if not timelines: + return {"solver_timeline_count": 0} + + ok_count = 0 + click_calls: list[float] = [] + click_move_ms: list[float] = [] + click_down_up_ms: list[float] = [] + render_to_token_ms: list[float] = [] + token_write_to_poll_done_ms: list[float] = [] + poll_attempts: list[float] = [] + poll_read_avg_ms: list[float] = [] + click_before_count = 0 + center_iframe_hits = 0 + turnstile_iframe_seen = 0 + widget_seen = 0 + + for timeline in timelines: + events = timeline.events + page_trace = _last_page_trace(events) + if page_trace: + render_called = page_trace.get("render_called_at") + token_written = page_trace.get("token_written_at") + if isinstance(render_called, (int, float)) and isinstance(token_written, (int, float)): + render_to_token_ms.append(float(token_written) - float(render_called)) + + for event in events: + if event.get("event") == "click_before": + click_before_count += 1 + dom = event.get("dom") if isinstance(event.get("dom"), dict) else {} + widget = dom.get("widget") if isinstance(dom.get("widget"), dict) else {} + center = dom.get("element_at_center") if isinstance(dom.get("element_at_center"), dict) else {} + if widget.get("present") and widget.get("visible"): + widget_seen += 1 + if center.get("is_iframe"): + center_iframe_hits += 1 + if (dom.get("turnstile_iframe_count") or 0) > 0: + turnstile_iframe_seen += 1 + + if event.get("event") == "click_after": + call_ms = event.get("click_call_ms") + if isinstance(call_ms, (int, float)): + click_calls.append(float(call_ms)) + trace = event.get("click_trace") if isinstance(event.get("click_trace"), dict) else {} + move1 = trace.get("mouse_move1_ms") + move2 = trace.get("mouse_move2_ms") + down = trace.get("mouse_down_ms") + up = trace.get("mouse_up_ms") + move_total = sum(float(v) for v in (move1, move2) if isinstance(v, (int, float))) + down_up_total = sum(float(v) for v in (down, up) if isinstance(v, (int, float))) + if move_total: + click_move_ms.append(move_total) + if down_up_total: + click_down_up_ms.append(down_up_total) + + if event.get("event") == "poll_done": + if event.get("ok"): + ok_count += 1 + attempts = event.get("poll_attempts") + if isinstance(attempts, (int, float)): + poll_attempts.append(float(attempts)) + read_avg = event.get("poll_read_ms_avg") + if isinstance(read_avg, (int, float)): + poll_read_avg_ms.append(float(read_avg)) + page_trace = event.get("page_trace") if isinstance(event.get("page_trace"), dict) else {} + token_written = page_trace.get("token_written_at") + event_t = event.get("t") + created_at = page_trace.get("created_at") + if all(isinstance(v, (int, float)) for v in (token_written, event_t, created_at)): + token_write_to_poll_done_ms.append((float(event_t) * 1000.0) - (float(token_written) - float(created_at))) + + return { + "solver_timeline_count": len(timelines), + "ok_count": ok_count, + "avg_click_call_ms": _avg(click_calls), + "avg_click_mouse_move_ms": _avg(click_move_ms), + "avg_click_down_up_ms": _avg(click_down_up_ms), + "avg_render_to_token_ms": _avg(render_to_token_ms), + "avg_token_write_to_poll_done_ms": _avg(token_write_to_poll_done_ms), + "avg_poll_attempts": _avg(poll_attempts), + "avg_poll_read_ms": _avg(poll_read_avg_ms), + "widget_seen_ratio": _ratio(widget_seen, click_before_count), + "center_iframe_hit_ratio": _ratio(center_iframe_hits, click_before_count), + "turnstile_iframe_seen_ratio": _ratio(turnstile_iframe_seen, click_before_count), + } + + +def analyze_text(text: str) -> dict[str, float | int | str | None]: + return summarize_monitor_rows(parse_monitor_lines(text)) diff --git a/xai_enroller/__init__.py b/xai_enroller/__init__.py new file mode 100644 index 0000000..3007474 --- /dev/null +++ b/xai_enroller/__init__.py @@ -0,0 +1,3 @@ +"""Standalone xAI OAuth Device Flow enroller.""" + +__version__ = "0.1.0" diff --git a/xai_enroller/__main__.py b/xai_enroller/__main__.py new file mode 100644 index 0000000..a9391dc --- /dev/null +++ b/xai_enroller/__main__.py @@ -0,0 +1,111 @@ +import argparse +import asyncio + +import httpx + +from .config import Settings +from .coordinator import EnrollmentCoordinator +from .executors import HTTPProbeExecutor, PlaywrightExecutor +from .protocol import XAIProfile, XAIProtocol +from .sinks import CPAAuthFileSink +from .sources import FileSourceAdapter, SQLiteSourceAdapter + + +def build_parser(): + parser = argparse.ArgumentParser(description="Bounded xAI OAuth Device Flow enroller") + parser.add_argument("--source", choices=("file", "sqlite")) + parser.add_argument("--target", type=int) + parser.add_argument("--concurrency", type=int) + parser.add_argument("--retry-attempts", type=int) + parser.add_argument("--executor", choices=("http", "playwright")) + parser.add_argument("--sink", choices=("cpa",)) + return parser + + +async def main_async(args=None): + parsed = build_parser().parse_args(args) + env = dict() + if parsed.source: + env["XAI_ENROLLER_SOURCE_KIND"] = parsed.source + if parsed.target is not None: + env["XAI_ENROLLER_TARGET"] = str(parsed.target) + if parsed.concurrency is not None: + env["XAI_ENROLLER_CONCURRENCY"] = str(parsed.concurrency) + if parsed.retry_attempts is not None: + env["XAI_ENROLLER_RETRY_ATTEMPTS"] = str(parsed.retry_attempts) + if parsed.executor: + env["XAI_ENROLLER_AUTH_EXECUTOR"] = parsed.executor + if parsed.sink: + env["XAI_ENROLLER_SINK"] = parsed.sink + import os + + merged = dict(os.environ) + merged.update(env) + settings = Settings.from_environ(merged) + if settings.source_kind == "remote": + raise ValueError("remote source requires the local auth service entrypoint") + source = ( + FileSourceAdapter(settings.source_file) + if settings.source_kind == "file" + else SQLiteSourceAdapter(settings.source_db, settings.source_salt) + ) + + # CF prewarm (external FlareSolverr) + register-side proxy for httpx/Playwright + try: + from grok_register.clearance import ( + format_prewarm_log, + httpx_proxy_mounts, + prewarm_clearance, + ) + + print(format_prewarm_log(prewarm_clearance())) + proxy_url = httpx_proxy_mounts() + except Exception as exc: # noqa: BLE001 + print(f"[clearance] skip: {exc}") + proxy_url = None + + client = httpx.AsyncClient(proxy=proxy_url) if proxy_url else httpx.AsyncClient() + protocol = XAIProtocol( + client, + XAIProfile.default(), + default_poll_interval=settings.poll_interval, + ) + executor = ( + HTTPProbeExecutor(client) + if settings.executor == "http" + else PlaywrightExecutor(settings.concurrency) + ) + sink = None + sink_client = None + if settings.sink == "cpa": + sink_client = ( + httpx.AsyncClient(proxy=proxy_url) if proxy_url else httpx.AsyncClient() + ) + sink = CPAAuthFileSink(settings.cpa_base_url, settings.cpa_management_secret, sink_client) + try: + coordinator = EnrollmentCoordinator( + source=source, + protocol=protocol, + executor=executor, + sink=sink, + ledger_path=settings.ledger_path, + ledger_salt=settings.source_salt, + concurrency=settings.concurrency, + timeout=settings.timeout_sec, + retry_attempts=settings.retry_attempts, + ) + results = await coordinator.run(settings.target) + for index, result in enumerate(results, 1): + print(f"job {index}: {result.status.value} ({result.reason_code})") + finally: + await client.aclose() + if sink_client: + await sink_client.aclose() + + +def main(): + asyncio.run(main_async()) + + +if __name__ == "__main__": + main() diff --git a/xai_enroller/auth_pipeline.py b/xai_enroller/auth_pipeline.py new file mode 100644 index 0000000..75d9c52 --- /dev/null +++ b/xai_enroller/auth_pipeline.py @@ -0,0 +1,1035 @@ +"""Bounded asynchronous local authentication pipeline.""" + +import asyncio +import heapq +import itertools +import time +from collections import deque +from contextlib import suppress +from dataclasses import replace + +import httpx + +from .models import ( + AuthorizationStatus, + CompletionJob, + JobStatus, + PipelineState, + PreparedJob, +) + + +class GlobalRateLimitGate: + """Process-wide cooldown with exactly one confirmation probe. + + Cooldown grows after consecutive trips so a cleared probe does not immediately + re-enter a tight authorization loop that re-trips the upstream limit. + """ + + COOLDOWN_SECONDS = 60.0 + MAX_COOLDOWN_SECONDS = 300.0 + COOLDOWN_GROWTH = 1.5 + + def __init__(self, *, clock=time.monotonic, base_cooldown=None): + self._clock = clock + self._condition = asyncio.Condition() + self._tripped_at = None + self._next_probe_at = 0.0 + self._probe_token = None + self._closed = False + self.base_cooldown = float( + self.COOLDOWN_SECONDS if base_cooldown is None else base_cooldown + ) + # Keep class attribute in sync for callers that still read COOLDOWN_SECONDS. + self.COOLDOWN_SECONDS = self.base_cooldown + self._current_cooldown = self.base_cooldown + self._consecutive_trips = 0 + + async def wait_for_permission(self, resume_event): + while True: + await resume_event.wait() + async with self._condition: + if self._closed: + raise asyncio.CancelledError + if not resume_event.is_set(): + continue + if self._tripped_at is None: + return None + now = self._clock() + if self._probe_token is None and now >= self._next_probe_at: + self._probe_token = object() + return self._probe_token + timeout = max(0.0, self._next_probe_at - now) + try: + if timeout: + await asyncio.wait_for(self._condition.wait(), timeout=timeout) + else: + await self._condition.wait() + except TimeoutError: + pass + + async def rate_limited(self, probe_token=None): + async with self._condition: + now = self._clock() + if self._tripped_at is None: + self._tripped_at = now + self._consecutive_trips = 1 + elif probe_token is None or probe_token is not self._probe_token: + return False + else: + self._consecutive_trips = max(1, self._consecutive_trips) + 1 + growth = self.COOLDOWN_GROWTH ** max(0, self._consecutive_trips - 1) + self._current_cooldown = min( + self.MAX_COOLDOWN_SECONDS, + self.base_cooldown * growth, + ) + self.COOLDOWN_SECONDS = self._current_cooldown + self._next_probe_at = now + self._current_cooldown + self._probe_token = None + self._condition.notify_all() + return True + + async def authorized(self, probe_token=None): + async with self._condition: + if self._tripped_at is None: + return None + if probe_token is None or probe_token is not self._probe_token: + return None + elapsed = max(0.0, self._clock() - self._tripped_at) + self._tripped_at = None + self._next_probe_at = 0.0 + self._probe_token = None + self._consecutive_trips = 0 + self._current_cooldown = self.base_cooldown + self.COOLDOWN_SECONDS = self.base_cooldown + self._condition.notify_all() + return elapsed + + async def inconclusive(self, probe_token=None): + async with self._condition: + if ( + self._tripped_at is None + or probe_token is None + or probe_token is not self._probe_token + ): + return False + self._next_probe_at = self._clock() + self._current_cooldown + self._probe_token = None + self._condition.notify_all() + return True + + async def close(self): + async with self._condition: + self._closed = True + self._probe_token = None + self._condition.notify_all() + + def snapshot(self): + now = self._clock() + return { + "cooldown": self._tripped_at is not None, + "probe_in_flight": self._probe_token is not None, + "cooldown_remaining_seconds": ( + round(max(0.0, self._next_probe_at - now), 1) + if self._tripped_at is not None + else 0.0 + ), + "cooldown_seconds": round(self._current_cooldown, 1), + "rate_limit_trips": self._consecutive_trips, + } + + +class MinimumStartInterval: + """Enforce a minimum interval between actual authorization starts. + + The configured base interval is the floor. After an upstream rate limit the + live interval steps up so the service does not immediately re-flood the + device-verify endpoint; it decays back toward the base after a streak of + successful authorizations. + """ + + RATE_LIMIT_FLOOR_SECONDS = 18.0 + RATE_LIMIT_GROWTH = 1.55 + SUCCESS_DECAY = 0.88 + SUCCESS_STREAK_TO_DECAY = 4 + + def __init__( + self, + seconds=0.0, + *, + max_seconds=None, + clock=time.monotonic, + sleep=asyncio.sleep, + ): + self.base_seconds = max(0.0, float(seconds)) + self.seconds = self.base_seconds + self.max_seconds = float( + max_seconds + if max_seconds is not None + else max(self.base_seconds * 8.0, 90.0) + ) + self._clock = clock + self._sleep = sleep + self._last_started_at = self._clock() if self.seconds > 0 else None + self._success_streak = 0 + + async def wait(self): + if self._last_started_at is None or self.seconds <= 0: + return + delay = self.seconds - (self._clock() - self._last_started_at) + if delay > 0: + await self._sleep(delay) + + def mark_started(self): + self._last_started_at = self._clock() + + def note_rate_limited(self): + """Raise the live start interval after an upstream rate-limit trip.""" + floor = max(self.base_seconds, self.RATE_LIMIT_FLOOR_SECONDS) + raised = max(self.seconds * self.RATE_LIMIT_GROWTH, floor) + self.seconds = min(self.max_seconds, raised) + self._success_streak = 0 + return self.seconds + + def note_success(self): + """Decay the live interval toward the configured base after successes.""" + self._success_streak += 1 + if ( + self._success_streak < self.SUCCESS_STREAK_TO_DECAY + or self.seconds <= self.base_seconds + ): + return self.seconds + self.seconds = max(self.base_seconds, self.seconds * self.SUCCESS_DECAY) + if self.seconds <= self.base_seconds + 0.05: + self.seconds = self.base_seconds + self._success_streak = 0 + return self.seconds + + def snapshot(self): + if self._last_started_at is None: + remaining = 0.0 + else: + remaining = max( + 0.0, + self.seconds - (self._clock() - self._last_started_at), + ) + return { + "min_authorization_interval_seconds": round(self.seconds, 1), + "base_authorization_interval_seconds": round(self.base_seconds, 1), + "pacing_remaining_seconds": round(remaining, 1), + } + + +class AuthPipeline: + SOURCE_QUEUE_CAPACITY = 64 + PREPARED_QUEUE_CAPACITY = 1 + COMPLETION_QUEUE_CAPACITY = 2 + FLOW_SAFETY_MARGIN_SECONDS = 60.0 + TRANSIENT_RETRY_DELAYS = (5.0, 15.0, 30.0) + CANCEL_RETRY_DELAY = 60.0 + + def __init__( + self, + *, + source, + protocol, + executor, + sink, + ledger, + timeout=1800.0, + min_authorization_interval=0.0, + event_callback=None, + clock=time.monotonic, + ): + self.source = source + self.protocol = protocol + self.executor = executor + self.sink = sink + self.ledger = ledger + self.timeout = float(timeout) + self.event_callback = event_callback + self._clock = clock + + self.source_queue = asyncio.Queue(maxsize=self.SOURCE_QUEUE_CAPACITY) + self.prepared_queue = asyncio.Queue(maxsize=self.PREPARED_QUEUE_CAPACITY) + self.completion_queue = asyncio.Queue(maxsize=self.COMPLETION_QUEUE_CAPACITY) + self.rate_gate = GlobalRateLimitGate(clock=clock) + self.start_interval = MinimumStartInterval( + min_authorization_interval, + clock=clock, + ) + + self._states = {} + self._transient_retries = {} + self._retry_heap = [] + self._retry_sequence = itertools.count() + self._retry_condition = asyncio.Condition() + self._resume = asyncio.Event() + self._resume.set() + self._prepare_slot_available = asyncio.Event() + self._prepare_slot_available.set() + self._stop_requested = asyncio.Event() + self._stopping = False + self._running = False + self._tasks = [] + + self._authorization_task = None + self._authorization_cancellable = False + self._completion_task = None + self._completion_cancellable = False + self._started_monotonic = self._clock() + self._process_imports = 0 + self._authorization_starts = 0 + self._recent_imports = deque() + + self.ledger.recover_pending() + for fingerprint in self.ledger.imported_fingerprints(): + self._states[fingerprint] = PipelineState.IMPORTED + + def _emit(self, kind, data): + if self.event_callback is None: + return + try: + self.event_callback(kind, data) + except Exception: + pass + + async def _guard(self, stage, coroutine): + try: + await coroutine + except asyncio.CancelledError: + raise + except Exception: + self._emit("pipeline_error", {"stage": stage, "reason": "internal_error"}) + self.request_stop() + + async def run(self): + if self._running: + raise RuntimeError("authentication pipeline is already running") + self._running = True + self._started_monotonic = self._clock() + try: + if hasattr(self.executor, "start"): + await self.executor.start() + self._tasks = [ + asyncio.create_task(self._guard("source", self._source_worker())), + asyncio.create_task(self._guard("retry", self._retry_worker())), + asyncio.create_task(self._guard("prepare", self._prepare_worker())), + asyncio.create_task( + self._guard("authorization", self._authorization_worker()) + ), + asyncio.create_task(self._guard("completion", self._completion_worker())), + ] + self._emit("service_started", self.start_interval.snapshot()) + await self._stop_requested.wait() + finally: + await self._shutdown() + + def request_stop(self): + self._stop_requested.set() + + async def _shutdown(self): + if self._stopping: + return + self._stopping = True + self._resume.set() + await self.rate_gate.close() + async with self._retry_condition: + self._retry_condition.notify_all() + for task in self._tasks: + task.cancel() + await asyncio.gather(*self._tasks, return_exceptions=True) + self._tasks.clear() + with suppress(Exception): + await self.source.close() + if hasattr(self.executor, "close"): + with suppress(Exception): + await self.executor.close() + self._settle_buffered_attempts() + self.ledger.recover_pending(reason="shutdown_cancelled") + self._running = False + self._emit("service_stopped", {}) + + def _settle_buffered_attempts(self): + for queue in (self.prepared_queue, self.completion_queue): + while True: + try: + item = queue.get_nowait() + except asyncio.QueueEmpty: + break + prepared = item.prepared if isinstance(item, CompletionJob) else item + self.ledger.finish( + prepared.job_id, JobStatus.CANCELLED, "shutdown_cancelled" + ) + queue.task_done() + while True: + try: + self.source_queue.get_nowait() + except asyncio.QueueEmpty: + break + self.source_queue.task_done() + self._retry_heap.clear() + + def pause(self): + self._resume.clear() + + def resume(self): + self._resume.set() + + async def cancel_active(self): + task = None + if ( + self._authorization_cancellable + and self._authorization_task is not None + and not self._authorization_task.done() + ): + task = self._authorization_task + self._authorization_cancellable = False + elif ( + self._completion_cancellable + and self._completion_task is not None + and not self._completion_task.done() + ): + task = self._completion_task + self._completion_cancellable = False + if task is None: + return False + task.cancel() + return True + + async def admit(self, source): + fingerprint = self.ledger.fingerprint(source.source_id) + if fingerprint in self._states: + return False + self._states[fingerprint] = PipelineState.QUEUED + try: + await self.source_queue.put(source) + except BaseException: + if self._states.get(fingerprint) is PipelineState.QUEUED: + self._states.pop(fingerprint, None) + raise + self._emit("source_admitted", {"queued": self.source_queue.qsize()}) + return True + + async def _source_worker(self): + async for source in self.source.records(): + if self._stopping: + return + await self.admit(source) + + def _set_state(self, fingerprint, state): + self._states[fingerprint] = state + + async def _schedule_retry(self, source, fingerprint, delay): + self._set_state(fingerprint, PipelineState.RETRY_WAITING) + due = self._clock() + delay + async with self._retry_condition: + heapq.heappush( + self._retry_heap, + (due, next(self._retry_sequence), fingerprint, source), + ) + self._retry_condition.notify_all() + + async def _retry_worker(self): + while not self._stopping: + async with self._retry_condition: + while not self._retry_heap and not self._stopping: + await self._retry_condition.wait() + if self._stopping: + return + due = self._retry_heap[0][0] + delay = due - self._clock() + if delay > 0: + try: + await asyncio.wait_for(self._retry_condition.wait(), timeout=delay) + except TimeoutError: + pass + continue + _due, _sequence, fingerprint, source = heapq.heappop(self._retry_heap) + if self._states.get(fingerprint) is not PipelineState.RETRY_WAITING: + continue + self._states[fingerprint] = PipelineState.QUEUED + await self.source_queue.put(source) + + async def _start_flow(self): + async with asyncio.timeout(self.timeout): + return await self.protocol.start_device_flow() + + async def _prepare_worker(self): + while not self._stopping: + await self._prepare_slot_available.wait() + self._prepare_slot_available.clear() + source = await self.source_queue.get() + fingerprint = self.ledger.fingerprint(source.source_id) + attempt = self.ledger.next_attempt(fingerprint) + job_id = self.ledger.start_fingerprint(fingerprint, attempt=attempt) + try: + try: + flow = await self._start_flow() + except asyncio.CancelledError: + self.ledger.finish(job_id, JobStatus.CANCELLED, "shutdown_cancelled") + raise + except ValueError: + await self._finish_failure( + source, + fingerprint, + job_id, + attempt, + JobStatus.SOURCE_INVALID, + "device_flow_invalid", + retry=False, + ) + self._prepare_slot_available.set() + continue + except Exception: + await self._finish_failure( + source, + fingerprint, + job_id, + attempt, + JobStatus.TRANSPORT_FAILED, + "device_flow_failed", + retry=True, + ) + self._prepare_slot_available.set() + continue + prepared = PreparedJob( + source=source, + source_fingerprint=fingerprint, + flow=flow, + flow_created_monotonic=self._clock(), + job_id=job_id, + attempt_number=attempt, + ) + self._set_state(fingerprint, PipelineState.PREPARED) + try: + await self.prepared_queue.put(prepared) + except asyncio.CancelledError: + self.ledger.finish(job_id, JobStatus.CANCELLED, "shutdown_cancelled") + raise + finally: + self.source_queue.task_done() + + def _flow_remaining(self, prepared): + age = max(0.0, self._clock() - prepared.flow_created_monotonic) + return float(prepared.flow.expires_in) - age + + async def _refresh_flow_if_needed(self, prepared): + if self._flow_remaining(prepared) >= self.FLOW_SAFETY_MARGIN_SECONDS: + return prepared + flow = await self._start_flow() + refreshed = replace( + prepared, + flow=flow, + flow_created_monotonic=self._clock(), + ) + if self._flow_remaining(refreshed) < self.FLOW_SAFETY_MARGIN_SECONDS: + raise RuntimeError("device flow lifetime below safety margin") + return refreshed + + async def _authorization_worker(self): + while not self._stopping: + prepared = await self.prepared_queue.get() + self._prepare_slot_available.set() + task = asyncio.create_task(self._authorize(prepared)) + self._authorization_task = task + try: + await task + except asyncio.CancelledError: + if self._stopping: + raise + finally: + self._authorization_cancellable = False + self._authorization_task = None + self.prepared_queue.task_done() + + async def _finish_if_already_imported(self, prepared, probe_token=None): + if not self.ledger.has_imported(prepared.source.source_id): + return False + self._authorization_cancellable = False + self.ledger.finish( + prepared.job_id, + JobStatus.CANCELLED, + "already_imported", + ) + self._set_state(prepared.source_fingerprint, PipelineState.IMPORTED) + self._transient_retries.pop(prepared.source_fingerprint, None) + await self.rate_gate.inconclusive(probe_token) + return True + + async def _authorize(self, prepared): + probe_token = None + self._authorization_cancellable = True + try: + await self.start_interval.wait() + await self._resume.wait() + if await self._finish_if_already_imported(prepared): + return + probe_token = await self.rate_gate.wait_for_permission(self._resume) + while True: + await self._resume.wait() + try: + prepared = await self._refresh_flow_if_needed(prepared) + except ValueError: + self._authorization_cancellable = False + await self.rate_gate.inconclusive(probe_token) + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.SOURCE_INVALID, + "device_flow_invalid", + retry=False, + ) + return + except Exception: + self._authorization_cancellable = False + await self.rate_gate.inconclusive(probe_token) + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.TRANSPORT_FAILED, + "device_flow_refresh_failed", + retry=True, + ) + return + if self._resume.is_set() and self._flow_remaining( + prepared + ) >= self.FLOW_SAFETY_MARGIN_SECONDS: + break + + if await self._finish_if_already_imported(prepared, probe_token): + return + + self._set_state(prepared.source_fingerprint, PipelineState.ACTIVE) + self.start_interval.mark_started() + self.ledger.mark_authorization_started(prepared.job_id) + self._authorization_starts += 1 + prepared = replace(prepared, task_number=self._authorization_starts) + self._emit( + "authorization_started", + { + "task_number": self._authorization_starts, + "attempt_number": prepared.attempt_number, + "source_queue": self.source_queue.qsize(), + "pending_total": self._pending_total(), + }, + ) + async with asyncio.timeout(self.timeout): + authorization = await self.executor.confirm( + prepared.source, prepared.flow + ) + self._authorization_cancellable = False + if authorization.status is AuthorizationStatus.AUTHORIZED: + recovery_elapsed = await self.rate_gate.authorized(probe_token) + if recovery_elapsed is not None: + self._emit( + "rate_limit_cleared", + {"elapsed_seconds": round(recovery_elapsed)}, + ) + paced = self.start_interval.note_success() + if paced > self.start_interval.base_seconds: + self._emit( + "pacing_adjusted", + { + "min_authorization_interval_seconds": round(paced, 1), + "reason": "success_decay", + }, + ) + await self.completion_queue.put(CompletionJob(prepared)) + return + + if authorization.reason_code == "rate_limited": + applied = await self.rate_gate.rate_limited(probe_token) + if applied: + paced = self.start_interval.note_rate_limited() + self._emit( + "rate_limited", + { + "wait_seconds": int(self.rate_gate.COOLDOWN_SECONDS), + "min_authorization_interval_seconds": round(paced, 1), + }, + ) + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.NEEDS_INTERACTION, + "rate_limited", + retry="rate_limited", + task_number=prepared.task_number, + ) + return + + await self.rate_gate.inconclusive(probe_token) + transient = authorization.reason_code in { + "browser_error", + "confirmation_timeout", + "challenge_required", + "unknown_page", + } + if authorization.reason_code == "confirmation_timeout": + status = JobStatus.TIMEOUT + else: + try: + status = JobStatus(authorization.status.value) + except ValueError: + status = JobStatus.NEEDS_INTERACTION + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + status, + authorization.reason_code, + retry=transient, + task_number=prepared.task_number, + ) + except asyncio.CancelledError: + self._authorization_cancellable = False + await self.rate_gate.inconclusive(probe_token) + reason = "shutdown_cancelled" if self._stopping else "operator_cancelled" + self.ledger.finish(prepared.job_id, JobStatus.CANCELLED, reason) + if not self._stopping: + await self._schedule_retry( + prepared.source, + prepared.source_fingerprint, + self.CANCEL_RETRY_DELAY, + ) + self._emit_result( + JobStatus.CANCELLED, + reason, + prepared.attempt_number, + prepared.task_number, + ) + raise + except TimeoutError: + self._authorization_cancellable = False + await self.rate_gate.inconclusive(probe_token) + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.TIMEOUT, + "confirmation_timeout", + retry=True, + task_number=prepared.task_number, + ) + except (httpx.TransportError, OSError): + self._authorization_cancellable = False + await self.rate_gate.inconclusive(probe_token) + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.TRANSPORT_FAILED, + "authorization_transport_failed", + retry=True, + task_number=prepared.task_number, + ) + except Exception: + self._authorization_cancellable = False + await self.rate_gate.inconclusive(probe_token) + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.NEEDS_INTERACTION, + "browser_error", + retry=True, + task_number=prepared.task_number, + ) + finally: + self._authorization_cancellable = False + + async def _completion_worker(self): + while not self._stopping: + completion = await self.completion_queue.get() + task = asyncio.create_task(self._complete(completion)) + self._completion_task = task + try: + await task + except asyncio.CancelledError: + if self._stopping: + raise + finally: + self._completion_cancellable = False + self._completion_task = None + self.completion_queue.task_done() + + async def _complete(self, completion): + prepared = completion.prepared + self._completion_cancellable = True + try: + async with asyncio.timeout(self.timeout): + credential = await self.protocol.poll_token( + endpoint=prepared.flow.token_endpoint, + flow=prepared.flow, + timeout=self.timeout, + ) + except asyncio.CancelledError: + self._completion_cancellable = False + await self._cancel_completion(prepared) + raise + except RuntimeError as error: + self._completion_cancellable = False + reason = str(error) + mapping = { + "oauth_denied": JobStatus.OAUTH_DENIED, + "oauth_expired": JobStatus.OAUTH_EXPIRED, + "oauth_rejected": JobStatus.OAUTH_REJECTED, + } + status = mapping.get(reason) + if status is not None: + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + status, + reason, + retry=False, + task_number=prepared.task_number, + ) + else: + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.TRANSPORT_FAILED, + "token_transport_failed", + retry=True, + task_number=prepared.task_number, + ) + return + except (httpx.TransportError, OSError, TimeoutError): + self._completion_cancellable = False + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.TRANSPORT_FAILED, + "token_transport_failed", + retry=True, + task_number=prepared.task_number, + ) + return + except Exception: + self._completion_cancellable = False + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.TRANSPORT_FAILED, + "token_failed", + retry=True, + task_number=prepared.task_number, + ) + return + + if self.sink is None: + self._completion_cancellable = False + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.SINK_FAILED, + "sink_unconfigured", + retry=False, + task_number=prepared.task_number, + ) + return + # Sink persistence and the matching ledger transition form one commit. + # LocalAuthFileSink uses a worker thread, so cancelling its await cannot + # cancel the underlying atomic write and would otherwise cause a duplicate + # authorization on retry. + self._completion_cancellable = False + store_task = asyncio.create_task(self.sink.store(credential)) + cancelled_during_commit = False + try: + try: + receipt = await asyncio.shield(store_task) + except asyncio.CancelledError: + cancelled_during_commit = True + receipt = await store_task + except Exception: + await self._finish_failure( + prepared.source, + prepared.source_fingerprint, + prepared.job_id, + prepared.attempt_number, + JobStatus.SINK_FAILED, + "sink_failed", + retry=True, + task_number=prepared.task_number, + ) + return + + self._completion_cancellable = False + self.ledger.finish( + prepared.job_id, + JobStatus.IMPORTED, + "imported", + receipt.fingerprint, + ) + self._set_state(prepared.source_fingerprint, PipelineState.IMPORTED) + self._transient_retries.pop(prepared.source_fingerprint, None) + self._process_imports += 1 + self._recent_imports.append(self._clock()) + self._emit_result( + JobStatus.IMPORTED, + "imported", + prepared.attempt_number, + prepared.task_number, + ) + if cancelled_during_commit: + raise asyncio.CancelledError + + async def _cancel_completion(self, prepared): + reason = "shutdown_cancelled" if self._stopping else "operator_cancelled" + self.ledger.finish(prepared.job_id, JobStatus.CANCELLED, reason) + if not self._stopping: + await self._schedule_retry( + prepared.source, + prepared.source_fingerprint, + self.CANCEL_RETRY_DELAY, + ) + self._emit_result( + JobStatus.CANCELLED, + reason, + prepared.attempt_number, + prepared.task_number, + ) + + async def _finish_failure( + self, + source, + fingerprint, + job_id, + attempt, + status, + reason, + *, + retry, + task_number=None, + ): + self.ledger.finish(job_id, status, reason) + if retry == "rate_limited": + await self._schedule_retry(source, fingerprint, 0.0) + elif retry: + retry_count = self._transient_retries.get(fingerprint, 0) + if retry_count < len(self.TRANSIENT_RETRY_DELAYS): + self._transient_retries[fingerprint] = retry_count + 1 + await self._schedule_retry( + source, + fingerprint, + self.TRANSIENT_RETRY_DELAYS[retry_count], + ) + else: + self._set_state(fingerprint, PipelineState.TERMINAL) + else: + self._set_state(fingerprint, PipelineState.TERMINAL) + self._emit_result(status, reason, attempt, task_number) + + def _prune_recent_imports(self): + cutoff = self._clock() - 300.0 + while self._recent_imports and self._recent_imports[0] < cutoff: + self._recent_imports.popleft() + + def _metrics(self): + self._prune_recent_imports() + counts = self.ledger.aggregate_counts() + elapsed = max(1.0, self._clock() - self._started_monotonic) + rolling_window = max(1.0, min(300.0, elapsed)) + finalized = counts["finalized_attempts"] + attempted_unique = counts["attempted_unique"] + return { + **counts, + **self.ledger.inventory_counts(), + "attempt_success": ( + counts["imported_attempts"] / finalized if finalized else 0.0 + ), + "eventual_success": ( + counts["imported_unique"] / attempted_unique + if attempted_unique + else 0.0 + ), + "five_minute_imports_per_minute": ( + len(self._recent_imports) * 60.0 / rolling_window + if self._process_imports + else None + ), + "lifetime_imports_per_minute": ( + self._process_imports * 60.0 / elapsed + if self._process_imports + else None + ), + "pending_total": self._pending_total(), + } + + def _pending_total(self): + snapshot = getattr(self.source, "snapshot_fingerprints", None) + if snapshot is None: + return None + return len(snapshot - self.ledger.imported_fingerprints()) + + def _emit_result(self, status, reason, attempt, task_number=None): + try: + metrics = self._metrics() + except Exception: + self._emit( + "pipeline_error", {"stage": "metrics", "reason": "internal_error"} + ) + self.request_stop() + return + self._emit( + "result", + { + "status": status.value, + "reason": reason, + "attempt_number": attempt, + "task_number": task_number, + **metrics, + }, + ) + + def status(self): + active_stages = [] + if self._authorization_task is not None and not self._authorization_task.done(): + active_stages.append("authorization") + if self._completion_task is not None and not self._completion_task.done(): + active_stages.append("completion") + retry_waiting = sum( + state is PipelineState.RETRY_WAITING for state in self._states.values() + ) + next_retry_seconds = ( + round(max(0.0, self._retry_heap[0][0] - self._clock()), 1) + if self._retry_heap + else 0.0 + ) + return { + "state": "stopping" if self._stopping else ( + "running" if self._resume.is_set() else "paused" + ), + "source_queue": self.source_queue.qsize(), + "prepared_queue": self.prepared_queue.qsize(), + "completion_queue": self.completion_queue.qsize(), + "retry_waiting": retry_waiting, + "next_retry_seconds": next_retry_seconds, + "active_stage": "+".join(active_stages) if active_stages else "idle", + "authorization_starts": self._authorization_starts, + **self.rate_gate.snapshot(), + **self.start_interval.snapshot(), + **self._metrics(), + } diff --git a/xai_enroller/config.py b/xai_enroller/config.py new file mode 100644 index 0000000..4357749 --- /dev/null +++ b/xai_enroller/config.py @@ -0,0 +1,139 @@ +import os +import stat +from dataclasses import dataclass +from pathlib import Path +from urllib.parse import urlparse + + +def _int(env, key, default): + try: + return int(env.get(key, str(default))) + except ValueError as exc: + raise ValueError(f"{key} must be an integer") from exc + + +def _float(env, key, default): + try: + return float(env.get(key, str(default))) + except ValueError as exc: + raise ValueError(f"{key} must be numeric") from exc + + +@dataclass(frozen=True) +class Settings: + source_kind: str + source_file: Path | None + source_db: Path | None + source_salt: bytes + ledger_path: Path + target: int = 1 + concurrency: int = 1 + retry_attempts: int = 0 + timeout_sec: float = 1800.0 + poll_interval: float = 5.0 + executor: str = "http" + sink: str | None = None + cpa_base_url: str | None = None + cpa_management_secret: str | None = None + local_auth_dir: Path | None = None + + @classmethod + def from_environ(cls, env=None): + env = dict(os.environ if env is None else env) + source_kind = env.get("XAI_ENROLLER_SOURCE_KIND", "file") + if source_kind not in {"file", "sqlite", "remote"}: + raise ValueError("source kind must be file, sqlite, or remote") + source_file = Path(env["XAI_ENROLLER_SOURCE_FILE"]) if env.get("XAI_ENROLLER_SOURCE_FILE") else None + source_db = Path(env["XAI_ENROLLER_SOURCE_DB"]) if env.get("XAI_ENROLLER_SOURCE_DB") else None + if source_kind == "file" and source_file is None: + raise ValueError("source file is required") + if source_kind == "sqlite" and source_db is None: + raise ValueError("source db is required") + salt = env.get("XAI_ENROLLER_SOURCE_SALT") + if not salt: + raise ValueError("source salt is required") + ledger_path = Path(env.get("XAI_ENROLLER_LEDGER_PATH", "xai-enroller-ledger.db")) + target = _int(env, "XAI_ENROLLER_TARGET", 1) + concurrency = _int(env, "XAI_ENROLLER_CONCURRENCY", 1) + retry_attempts = _int(env, "XAI_ENROLLER_RETRY_ATTEMPTS", 0) + timeout_sec = _float(env, "XAI_ENROLLER_TIMEOUT_SEC", 1800) + poll_interval = _float(env, "XAI_ENROLLER_POLL_SEC", 5) + executor = env.get("XAI_ENROLLER_AUTH_EXECUTOR", "http") + sink = env.get("XAI_ENROLLER_SINK") or None + if not 1 <= target <= 100: + raise ValueError("target must be between 1 and 100") + if not 1 <= concurrency <= 4: + raise ValueError("concurrency must be between 1 and 4") + if not 0 <= retry_attempts <= 3: + raise ValueError("retry attempts must be between 0 and 3") + if timeout_sec <= 0: + raise ValueError("timeout must be positive") + if poll_interval <= 0: + raise ValueError("poll interval must be positive") + if executor not in {"http", "playwright"}: + raise ValueError("executor must be http or playwright") + if sink not in {None, "cpa", "local"}: + raise ValueError("sink must be cpa or local") + cpa_url = env.get("XAI_ENROLLER_CPA_BASE_URL") or None + cpa_secret = env.get("XAI_ENROLLER_CPA_MANAGEMENT_SECRET") or None + local_auth_dir = ( + Path(env["XAI_ENROLLER_LOCAL_AUTH_DIR"]).expanduser() + if env.get("XAI_ENROLLER_LOCAL_AUTH_DIR") + else None + ) + if sink == "cpa": + if not cpa_url or not cpa_secret: + raise ValueError("CPA base URL and management secret are required") + parsed = urlparse(cpa_url) + is_private = parsed.hostname in {"localhost", "127.0.0.1", "::1"} + if parsed.scheme != "https" and not is_private: + raise ValueError("CPA base URL must use HTTPS") + if sink == "local" and local_auth_dir is None: + raise ValueError("local auth directory is required") + return cls( + source_kind=source_kind, + source_file=source_file, + source_db=source_db, + source_salt=salt.encode(), + ledger_path=ledger_path, + target=target, + concurrency=concurrency, + retry_attempts=retry_attempts, + timeout_sec=timeout_sec, + poll_interval=poll_interval, + executor=executor, + sink=sink, + cpa_base_url=cpa_url, + cpa_management_secret=cpa_secret, + local_auth_dir=local_auth_dir, + ) + + def redacted_dict(self): + return { + "source_kind": self.source_kind, + "source_file": str(self.source_file) if self.source_file else None, + "source_db": str(self.source_db) if self.source_db else None, + "ledger_path": str(self.ledger_path), + "target": self.target, + "concurrency": self.concurrency, + "retry_attempts": self.retry_attempts, + "timeout_sec": self.timeout_sec, + "poll_interval": self.poll_interval, + "executor": self.executor, + "sink": self.sink, + "cpa_base_url": self.cpa_base_url, + "local_auth_dir": str(self.local_auth_dir) if self.local_auth_dir else None, + } + + +def require_private_regular_file(path: Path, *, require_0600=False): + try: + mode = path.stat() + except OSError as exc: + raise ValueError(f"source path is unavailable: {path}") from exc + if not stat.S_ISREG(mode.st_mode): + raise ValueError("source path must be a regular file") + if mode.st_mode & (stat.S_IWGRP | stat.S_IWOTH): + raise ValueError("source path must not be group/world writable") + if require_0600 and mode.st_mode & (stat.S_IRGRP | stat.S_IROTH): + raise ValueError("file source requires mode 0600 or stricter") diff --git a/xai_enroller/coordinator.py b/xai_enroller/coordinator.py new file mode 100644 index 0000000..5eb02e6 --- /dev/null +++ b/xai_enroller/coordinator.py @@ -0,0 +1,177 @@ +import asyncio +from contextlib import suppress + +import httpx + +from .ledger import Ledger +from .models import ( + AuthorizationStatus, + JobResult, + JobStatus, +) + + +class EnrollmentCoordinator: + def __init__( + self, + *, + source, + protocol, + executor, + sink, + ledger_path, + ledger_salt, + concurrency=1, + timeout=1800, + retry_attempts=0, + event_callback=None, + ): + self.source = source + self.protocol = protocol + self.executor = executor + self.sink = sink + self.ledger = Ledger(ledger_path, ledger_salt) + self.concurrency = concurrency + self.timeout = timeout + self.retry_attempts = max(0, retry_attempts) + self.event_callback = event_callback + self._semaphore = asyncio.Semaphore(concurrency) + + @staticmethod + def _is_pre_device_transport_failure(error): + return isinstance(error, (httpx.TransportError, OSError)) + + def _emit(self, kind, data): + if self.event_callback is None: + return + try: + self.event_callback(kind, data) + except Exception: + pass + + async def run(self, target=1): + return await self.run_records(self.source.records(), target=target) + + async def run_records(self, records, target=100): + if not 1 <= target <= 100: + raise ValueError("target must be between 1 and 100") + self.ledger.recover_pending() + results = [] + tasks = [] + index = 0 + if hasattr(records, "__aiter__"): + async for source in records: + if index >= target: + break + index += 1 + tasks.append(asyncio.create_task(self._run_one(source))) + else: + for source in records: + if index >= target: + break + index += 1 + tasks.append(asyncio.create_task(self._run_one(source))) + if hasattr(self.executor, "start"): + await self.executor.start() + try: + if tasks: + results = await asyncio.gather(*tasks) + return results + finally: + for task in tasks: + if not task.done(): + task.cancel() + if hasattr(self.executor, "close"): + with suppress(Exception): + await self.executor.close() + + async def _run_one(self, source): + async with self._semaphore: + for attempt in range(1, self.retry_attempts + 2): + job_id = self.ledger.start(source.source_id, attempt=attempt) + try: + async with asyncio.timeout(self.timeout): + try: + flow = await self.protocol.start_device_flow() + except Exception as error: + result = JobResult( + source.source_id, + JobStatus.TRANSPORT_FAILED, + "device_flow_failed", + attempt, + ) + self.ledger.finish(job_id, result.status, result.reason_code) + if ( + self._is_pre_device_transport_failure(error) + and attempt <= self.retry_attempts + ): + continue + return result + self._emit( + "device_flow", + { + "source_id": source.source_id, + "user_code": flow.user_code, + "verification_url": flow.verification_url, + }, + ) + return await self._attempt_after_flow(source, job_id, flow, attempt) + except asyncio.TimeoutError: + result = JobResult(source.source_id, JobStatus.TIMEOUT, "timeout", attempt) + self.ledger.finish(job_id, result.status, result.reason_code) + return result + except asyncio.CancelledError: + result = JobResult(source.source_id, JobStatus.CANCELLED, "cancelled", attempt) + self.ledger.finish(job_id, result.status, result.reason_code) + raise + except Exception: + result = JobResult( + source.source_id, JobStatus.TRANSPORT_FAILED, "transport_failed", attempt + ) + self.ledger.finish(job_id, result.status, result.reason_code) + return result + + async def _attempt_after_flow(self, source, job_id, flow, attempt): + authorization = await self.executor.confirm(source, flow) + if authorization.status is not AuthorizationStatus.AUTHORIZED: + status = JobStatus(authorization.status.value) + result = JobResult(source.source_id, status, authorization.reason_code, attempt) + self.ledger.finish(job_id, result.status, result.reason_code) + return result + if self.sink is None: + result = JobResult(source.source_id, JobStatus.SINK_FAILED, "sink_unconfigured", attempt) + self.ledger.finish(job_id, result.status, result.reason_code) + return result + try: + credential = await self.protocol.poll_token( + endpoint=flow.token_endpoint, + flow=flow, + timeout=self.timeout, + ) + except RuntimeError as error: + reason = str(error) + mapping = { + "oauth_denied": JobStatus.OAUTH_DENIED, + "oauth_expired": JobStatus.OAUTH_EXPIRED, + "oauth_rejected": JobStatus.OAUTH_REJECTED, + } + result = JobResult( + source.source_id, mapping.get(reason, JobStatus.TRANSPORT_FAILED), reason, attempt + ) + self.ledger.finish(job_id, result.status, result.reason_code) + return result + try: + receipt = await self.sink.store(credential) + except Exception: + result = JobResult(source.source_id, JobStatus.SINK_FAILED, "sink_failed", attempt) + self.ledger.finish(job_id, result.status, result.reason_code) + return result + result = JobResult( + source.source_id, + JobStatus.IMPORTED, + "imported", + attempt, + sink_receipt_fingerprint=receipt.fingerprint, + ) + self.ledger.finish(job_id, result.status, result.reason_code, receipt.fingerprint) + return result diff --git a/xai_enroller/executors.py b/xai_enroller/executors.py new file mode 100644 index 0000000..066a415 --- /dev/null +++ b/xai_enroller/executors.py @@ -0,0 +1,595 @@ +import asyncio +import glob +import importlib +import os +from contextlib import suppress +from urllib.parse import parse_qs, urlparse + +import httpx + +from grok_register.clearance import ( + apply_clearance_to_context, + cached_user_agent, + playwright_proxy_settings, +) + +from .models import AuthorizationResult, AuthorizationStatus, DeviceFlow, SourceRecord + + +class HTTPProbeExecutor: + def __init__(self, client: httpx.AsyncClient, max_body=64 * 1024): + self.client = client + self.max_body = max_body + + async def confirm(self, source: SourceRecord, flow: DeviceFlow): + parsed = urlparse(flow.verification_url) + if parsed.scheme != "https" or parsed.hostname != "accounts.x.ai": + return AuthorizationResult(AuthorizationStatus.NEEDS_INTERACTION, "invalid_url") + response = await self.client.get( + flow.verification_url, + headers={"Cookie": f"sso={source.sso_token}"}, + follow_redirects=False, + ) + body = await response.aread() + if len(body) > self.max_body: + return AuthorizationResult(AuthorizationStatus.NEEDS_INTERACTION, "response_too_large") + text = body.decode("utf-8", errors="replace").lower() + if response.status_code == 403 and any( + marker in text for marker in ("challenge", "cf-chl", "captcha") + ): + return AuthorizationResult(AuthorizationStatus.NEEDS_BROWSER, "challenge") + if response.status_code in {301, 302, 303, 307, 308}: + return AuthorizationResult(AuthorizationStatus.NEEDS_INTERACTION, "redirect") + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, "http_probe_inconclusive" + ) + + +class PlaywrightExecutor: + ALLOWED_CONTROLS = frozenset({"Authorize", "Allow", "Continue", "Confirm"}) + ATTEMPT_TIMEOUT_SECONDS = 55.0 + CLOSE_TIMEOUT_SECONDS = 5.0 + # Keep confirm loops tight: most successful device verifies finish under ~15s. + LOOP_DEADLINE_SECONDS = 28.0 + COOKIE_WAIT_MS = 250 + CONSENT_WAIT_MS = 350 + CODE_WAIT_MS = 300 + CHALLENGE_WAIT_MS = 700 + POLL_WAIT_MS = 250 + # Drop heavy assets that never participate in device verification. + BLOCKED_RESOURCE_TYPES = frozenset({"image", "media", "font", "stylesheet"}) + BLOCKED_URL_MARKERS = ( + "google-analytics", + "googletagmanager", + "doubleclick", + "facebook.net", + "hotjar", + "segment.io", + ) + + def __init__(self, concurrency=1, playwright_factory=None, executable_path=None): + self.concurrency = concurrency + self.playwright_factory = playwright_factory + self.executable_path = executable_path or self._find_executable_path() + self._playwright = None + self._browser = None + self._semaphore = asyncio.Semaphore(concurrency) + self._lifecycle_lock = asyncio.Lock() + + @staticmethod + def _find_executable_path(): + configured = os.environ.get("XAI_ENROLLER_BROWSER_EXECUTABLE") + if configured: + return configured + candidates = glob.glob(os.path.expanduser("~/.cloakbrowser/chromium-*/chrome")) + candidates.extend( + glob.glob( + os.path.expanduser( + "~/.cloakbrowser/chromium-*/Chromium.app/Contents/MacOS/Chromium" + ) + ) + ) + return sorted(candidates)[-1] if candidates else None + + async def start(self): + async with self._lifecycle_lock: + if self._browser is not None: + is_connected = getattr(self._browser, "is_connected", None) + if not callable(is_connected) or is_connected(): + return + await self._close_unlocked() + factory = self.playwright_factory + if factory is None: + module = importlib.import_module("playwright.async_api") + factory = module.async_playwright + playwright = await factory().start() + options = {"headless": True} + if self.executable_path: + options["executable_path"] = self.executable_path + # REGISTER_PROXY / HTTP_PROXY / HTTPS_PROXY → this process only + proxy = playwright_proxy_settings() + if proxy: + options["proxy"] = proxy + try: + browser = await playwright.chromium.launch(**options) + except BaseException: + await self._close_transport_resource(playwright.stop()) + raise + self._playwright = playwright + self._browser = browser + + async def close(self): + async with self._lifecycle_lock: + await self._close_unlocked() + + async def _close_unlocked(self): + browser = self._browser + playwright = self._playwright + self._browser = None + self._playwright = None + try: + if browser is not None: + await self._close_transport_resource(browser.close()) + finally: + if playwright is not None: + await self._close_transport_resource(playwright.stop()) + + @staticmethod + def _consume_task_result(task): + with suppress(BaseException): + task.result() + + @classmethod + async def _close_transport_resource(cls, awaitable): + task = asyncio.create_task(awaitable) + try: + done, _pending = await asyncio.wait( + {task}, timeout=cls.CLOSE_TIMEOUT_SECONDS + ) + except asyncio.CancelledError: + task.cancel() + task.add_done_callback(cls._consume_task_result) + raise + if task in done: + cls._consume_task_result(task) + return True + task.cancel() + task.add_done_callback(cls._consume_task_result) + return False + + async def _recycle_browser(self, browser): + async with self._lifecycle_lock: + if self._browser is browser: + await self._close_unlocked() + + async def __aenter__(self): + await self.start() + return self + + async def __aexit__(self, exc_type, exc, tb): + await self.close() + + @classmethod + async def _close_attempt_resources(cls, page, context): + async def close_session(): + closers = [] + if page is not None: + closers.append(page.close()) + if context is not None: + closers.append(context.close()) + if closers: + await asyncio.gather(*closers, return_exceptions=True) + + cleanup = asyncio.create_task(close_session()) + cancelled = False + deadline = asyncio.get_running_loop().time() + cls.CLOSE_TIMEOUT_SECONDS + while not cleanup.done(): + remaining = deadline - asyncio.get_running_loop().time() + if remaining <= 0: + break + try: + await asyncio.wait({cleanup}, timeout=remaining) + except asyncio.CancelledError: + cancelled = True + if cleanup.done(): + cls._consume_task_result(cleanup) + else: + # asyncio.wait_for() is not a hard deadline: it cancels the child and + # then waits for cancellation to finish, which a wedged Playwright + # transport may never do. Abandon the cleanup task after one real + # wall-clock deadline and consume its eventual result asynchronously. + cleanup.cancel() + cleanup.add_done_callback(cls._consume_task_result) + if cancelled: + raise asyncio.CancelledError + return cleanup.done() + + @staticmethod + async def _click_visible_exact(page, names): + for name in names: + try: + button = page.get_by_role("button", name=name, exact=True) + except TypeError: + button = page.get_by_role("button", name=name) + for index in range(await button.count()): + candidate = button.nth(index) if hasattr(button, "nth") else button.first + is_visible = getattr(candidate, "is_visible", None) + if is_visible is None or await is_visible(): + await candidate.click() + return True + return False + + @staticmethod + async def _click_visible_turnstile(page): + """Click a visible Turnstile widget using the registration solver motion.""" + for selector in ( + ".cf-turnstile", + "iframe[src*='challenges.cloudflare.com']", + "iframe[src*='turnstile']", + ): + locator = page.locator(selector).first + if not await locator.count() or not await locator.is_visible(): + continue + box = await locator.bounding_box() + if not box: + continue + x = box["x"] + box["width"] / 2 + y = box["y"] + box["height"] / 2 + await page.mouse.move(max(0, x - 25), max(0, y - 8)) + await page.mouse.move(x, y, steps=8) + await page.mouse.down() + await asyncio.sleep(0.05) + await page.mouse.up() + return True + return False + + @staticmethod + def _expanded_cookies(source): + allowed = { + "name", + "value", + "url", + "domain", + "path", + "expires", + "httpOnly", + "secure", + "sameSite", + } + cookies = [] + for source_cookie in source.cookies: + cookie = { + key: value for key, value in source_cookie.items() if key in allowed + } + if cookie.get("domain"): + cookie.pop("url", None) + elif cookie.get("url"): + cookie.pop("domain", None) + cookie.pop("path", None) + cookies.append(cookie) + if not cookies: + cookies = [ + { + "name": "sso", + "value": source.sso_token, + "domain": "accounts.x.ai", + "path": "/", + "secure": True, + "httpOnly": True, + } + ] + seen = { + (cookie.get("name"), cookie.get("domain"), cookie.get("path", "/")) + for cookie in cookies + } + for cookie in list(cookies) if source.cookies else (): + name = cookie.get("name", "") + if not name.startswith("sso"): + continue + key = (name, ".x.ai", cookie.get("path", "/")) + if key in seen: + continue + clone = dict(cookie) + clone.pop("url", None) + clone["domain"] = ".x.ai" + clone.setdefault("path", "/") + cookies.append(clone) + seen.add(key) + return cookies + + async def confirm(self, source: SourceRecord, flow: DeviceFlow): + await self.start() + async with self._semaphore: + browser = self._browser + attempt = asyncio.create_task( + self._confirm_once(browser, source, flow) + ) + try: + done, _pending = await asyncio.wait( + {attempt}, timeout=self.ATTEMPT_TIMEOUT_SECONDS + ) + except asyncio.CancelledError: + attempt.cancel() + attempt.add_done_callback(self._consume_task_result) + raise + if attempt in done: + return attempt.result() + + # A Playwright transport can swallow cancellation while one of its + # RPCs is wedged. Do not let asyncio.wait_for() extend the nominal + # timeout indefinitely: detach the attempt and replace the complete + # browser transport before admitting the next account. + attempt.cancel() + attempt.add_done_callback(self._consume_task_result) + await self._recycle_browser(browser) + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, "confirmation_timeout" + ) + + async def _install_request_filters(self, page): + """Skip non-essential resources so device verification spends less wall time.""" + + async def route_handler(route): + request = route.request + resource_type = getattr(request, "resource_type", "") or "" + url = getattr(request, "url", "") or "" + if resource_type in self.BLOCKED_RESOURCE_TYPES or any( + marker in url for marker in self.BLOCKED_URL_MARKERS + ): + await route.abort() + return + await route.continue_() + + with suppress(Exception): + await page.route("**/*", route_handler) + + async def _confirm_once(self, browser, source: SourceRecord, flow: DeviceFlow): + context = None + page = None + code_submitted = False + consent_submitted = False + challenge_clicks = 0 + try: + context_kwargs = {} + ua = cached_user_agent() + if ua: + context_kwargs["user_agent"] = ua + context = await browser.new_context(**context_kwargs) + # CF cookies from external FlareSolverr prewarm (host-level) + await apply_clearance_to_context(context) + page = await context.new_page() + await self._install_request_filters(page) + await context.add_cookies(self._expanded_cookies(source)) + await page.goto(flow.verification_url, wait_until="domcontentloaded") + deadline = ( + asyncio.get_running_loop().time() + self.LOOP_DEADLINE_SECONDS + ) + unknown_since = None + while asyncio.get_running_loop().time() < deadline: + parsed = urlparse(page.url) + if parsed.scheme != "https" or not parsed.hostname or not ( + parsed.hostname == "x.ai" + or parsed.hostname.endswith(".x.ai") + ): + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, "unsafe_page" + ) + query_error = parse_qs(parsed.query).get("error", [None])[0] + text = await page.locator("body").inner_text() + text_lower = text.lower() + page_title = getattr(page, "title", None) + title_lower = ( + (await page_title()).lower() if page_title is not None else "" + ) + + if query_error == "rate_limited": + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, "rate_limited" + ) + if query_error: + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, "device_verify_failed" + ) + if any( + marker in text_lower + for marker in ("rate limit", "too many requests", "请求过于频繁") + ): + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, "rate_limited" + ) + if ( + "attention required" in title_lower + and "cloudflare" in title_lower + ) or any( + marker in text_lower + for marker in ( + "sorry, you have been blocked", + "unable to access x.ai", + "cloudflare ray id", + ) + ): + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, + "challenge_required", + ) + if "/oauth2/device/done" in parsed.path or any( + marker in text_lower + for marker in ("device authorized", "设备已授权") + ): + return AuthorizationResult( + AuthorizationStatus.AUTHORIZED, "confirmed" + ) + + cookie_clicked = await self._click_visible_exact( + page, + ( + "全部拒绝", + "拒绝全部", + "Reject all", + "Reject All", + ), + ) + if cookie_clicked: + unknown_since = None + await page.wait_for_timeout(self.COOKIE_WAIT_MS) + continue + + if "/oauth2/device/consent" in parsed.path or any( + marker in text_lower + for marker in ("authorize grok build", "授权 grok build") + ): + if not consent_submitted: + allowed = await self._click_visible_exact( + page, ("允许", "Allow", "Authorize", "Approve") + ) + if allowed: + consent_submitted = True + unknown_since = None + await page.wait_for_timeout(self.CONSENT_WAIT_MS) + continue + else: + await page.wait_for_timeout(self.POLL_WAIT_MS) + continue + + code_input = page.locator('input[name="user_code"]') + try: + code_input_count = await code_input.count() + except AttributeError: + if await self._click_visible_exact( + page, self.ALLOWED_CONTROLS + ): + return AuthorizationResult( + AuthorizationStatus.AUTHORIZED, + "confirmed", + ) + code_input_count = 0 + if code_input_count and await code_input.first.is_visible(): + if code_submitted: + await page.wait_for_timeout(self.POLL_WAIT_MS) + continue + unknown_since = None + current = await code_input.first.input_value() + # fill() is far cheaper than press_sequentially and the + # device-code form does not require keystroke events. + if flow.user_code.replace("-", "") not in current.replace( + "-", "" + ): + await code_input.first.fill(flow.user_code) + submit = page.locator( + 'button[type="submit"], input[type="submit"]' + ) + if not await submit.count(): + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, + "submit_missing", + ) + predicate = lambda response: ( + urlparse(response.url).hostname == "auth.x.ai" + and urlparse(response.url).path == "/oauth2/device/verify" + ) + async with page.expect_response( + predicate, timeout=12000 + ) as response_info: + code_submitted = True + await submit.first.click() + response = await response_info.value + location = urlparse(response.headers.get("location", "")) + error = parse_qs(location.query).get("error", [None])[0] + if error == "rate_limited": + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, + "rate_limited", + ) + if error: + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, + "device_verify_failed", + ) + await page.wait_for_timeout(self.CODE_WAIT_MS) + continue + + if any( + marker in text_lower + for marker in ( + "continue with email", + "sign in with email", + "使用邮箱登录", + ) + ) or await page.locator('input[type="password"]').count(): + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, + "login_required", + ) + + if any( + marker in text_lower + for marker in ( + "captcha", + "verify you are human", + "security challenge", + "turnstile", + "人机验证", + "安全验证", + ) + ): + if ( + challenge_clicks < 3 + and await self._click_visible_turnstile(page) + ): + challenge_clicks += 1 + unknown_since = None + await page.wait_for_timeout(self.CHALLENGE_WAIT_MS) + continue + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, + "challenge_required", + ) + + if any( + marker in text_lower + for marker in ( + "multi-factor", + "two-factor", + "authentication code", + "verification code", + "双重验证", + "验证码", + ) + ): + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, + "mfa_required", + ) + + if await self._click_visible_exact(page, self.ALLOWED_CONTROLS): + return AuthorizationResult( + AuthorizationStatus.AUTHORIZED, + "confirmed", + ) + + now = asyncio.get_running_loop().time() + if unknown_since is None: + unknown_since = now + elif now - unknown_since >= 2.5: + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, "unknown_page" + ) + + await page.wait_for_timeout(self.POLL_WAIT_MS) + + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, "confirmation_timeout" + ) + except asyncio.CancelledError: + raise + except Exception as error: + is_timeout = ( + error.__class__.__name__ == "TimeoutError" + and error.__class__.__module__.startswith("playwright") + ) + return AuthorizationResult( + AuthorizationStatus.NEEDS_INTERACTION, + "confirmation_timeout" if is_timeout else "browser_error", + ) + finally: + await self._close_attempt_resources(page, context) diff --git a/xai_enroller/inventory.py b/xai_enroller/inventory.py new file mode 100644 index 0000000..75bc432 --- /dev/null +++ b/xai_enroller/inventory.py @@ -0,0 +1,105 @@ +import os +import re +import secrets +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path + + +class InventoryError(RuntimeError): + pass + + +@dataclass(frozen=True) +class ClaimBatch: + batch_id: str + directory: Path + moved: int + note: str = "" + + +class CredentialInventory: + def __init__(self, ledger, available_directory, claimed_directory): + self.ledger = ledger + self.available_directory = Path(available_directory) + self.claimed_directory = Path(claimed_directory) + + def take(self, limit, note=""): + try: + limit = int(limit) + except (TypeError, ValueError) as exc: + raise ValueError("claim count must be an integer") from exc + if limit <= 0: + raise ValueError("claim count must be positive") + note = str(note or "").strip() + batch_id = self._new_batch_id() + receipts = self.ledger.claim_available(limit, batch_id) + return self._complete_batch(batch_id, receipts, note) + + def recover(self): + rows = self.ledger.pending_claims() + batches = {} + for row in rows: + batches.setdefault(row["batch_id"], []).append( + row["sink_receipt_fingerprint"] + ) + recovered = 0 + for batch_id, receipts in batches.items(): + recovered += self._complete_batch(batch_id, receipts, "").moved + return recovered + + @staticmethod + def _new_batch_id(): + timestamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") + return f"{timestamp}-{secrets.token_hex(3)}" + + @staticmethod + def _validate_receipt(receipt): + if not re.fullmatch(r"[A-Za-z0-9][A-Za-z0-9._-]{0,127}", receipt): + raise InventoryError("invalid credential receipt") + + @staticmethod + def _fsync_directory(directory): + descriptor = os.open(directory, os.O_RDONLY) + try: + os.fsync(descriptor) + finally: + os.close(descriptor) + + def _complete_batch(self, batch_id, receipts, note): + destination_directory = self.claimed_directory / batch_id + if not receipts: + return ClaimBatch(batch_id, destination_directory, 0, note) + self.available_directory.mkdir(mode=0o700, parents=True, exist_ok=True) + self.claimed_directory.mkdir(mode=0o700, parents=True, exist_ok=True) + destination_directory.mkdir(mode=0o700, parents=False, exist_ok=True) + for directory in ( + self.available_directory, + self.claimed_directory, + destination_directory, + ): + os.chmod(directory, 0o700) + + moved = 0 + for receipt in receipts: + self._validate_receipt(receipt) + source = self.available_directory / f"{receipt}.json" + destination = destination_directory / source.name + source_exists = source.is_file() + destination_exists = destination.is_file() + if source_exists and destination_exists: + raise InventoryError("credential exists in both inventory locations") + if not source_exists and not destination_exists: + raise InventoryError("credential file is missing") + if source_exists: + os.replace(source, destination) + os.chmod(destination, 0o600) + moved += 1 + + self._fsync_directory(self.available_directory) + self._fsync_directory(destination_directory) + self._fsync_directory(self.claimed_directory) + marked = self.ledger.mark_claimed(batch_id, note=note) + if marked != len(receipts): + raise InventoryError("inventory batch did not commit completely") + return ClaimBatch(batch_id, destination_directory, moved, note) diff --git a/xai_enroller/ledger.py b/xai_enroller/ledger.py new file mode 100644 index 0000000..19996ac --- /dev/null +++ b/xai_enroller/ledger.py @@ -0,0 +1,275 @@ +import hashlib +import hmac +import os +import sqlite3 +from datetime import datetime, timezone +from pathlib import Path + +from .models import JobStatus + + +class Ledger: + def __init__(self, path: Path, salt: bytes): + self.path = Path(path) + self.salt = bytes(salt) + self._init() + + def _connect(self): + connection = sqlite3.connect(self.path) + connection.row_factory = sqlite3.Row + return connection + + def _init(self): + with self._connect() as connection: + connection.execute( + """ + CREATE TABLE IF NOT EXISTS jobs ( + job_id INTEGER PRIMARY KEY, + source_fingerprint TEXT NOT NULL, + attempt_number INTEGER NOT NULL, + status TEXT NOT NULL, + started_at TEXT NOT NULL, + finished_at TEXT, + reason_code TEXT, + sink_receipt_fingerprint TEXT + ) + """ + ) + columns = { + row["name"] for row in connection.execute("PRAGMA table_info(jobs)") + } + if "authorization_started" not in columns: + connection.execute( + "ALTER TABLE jobs ADD COLUMN authorization_started INTEGER NOT NULL DEFAULT 0" + ) + connection.execute( + "CREATE INDEX IF NOT EXISTS jobs_source_status_idx " + "ON jobs(source_fingerprint, status)" + ) + connection.execute( + """ + CREATE TABLE IF NOT EXISTS credential_inventory ( + sink_receipt_fingerprint TEXT PRIMARY KEY, + state TEXT NOT NULL CHECK(state IN ('available', 'claiming', 'claimed')), + claimed_at TEXT NOT NULL DEFAULT '', + batch_id TEXT NOT NULL DEFAULT '', + note TEXT NOT NULL DEFAULT '' + ) + """ + ) + connection.execute( + "CREATE INDEX IF NOT EXISTS credential_inventory_state_idx " + "ON credential_inventory(state)" + ) + connection.execute( + """ + INSERT OR IGNORE INTO credential_inventory( + sink_receipt_fingerprint, state + ) + SELECT sink_receipt_fingerprint, 'available' + FROM jobs + WHERE status=? + AND COALESCE(TRIM(sink_receipt_fingerprint), '') <> '' + """, + (JobStatus.IMPORTED.value,), + ) + os.chmod(self.path, 0o600) + + def fingerprint(self, source_id): + return hmac.new(self.salt, source_id.encode(), hashlib.sha256).hexdigest() + + def _fingerprint(self, source_id): + return self.fingerprint(source_id) + + def start(self, source_id, attempt=1): + return self.start_fingerprint(self.fingerprint(source_id), attempt=attempt) + + def start_fingerprint(self, source_fingerprint, attempt=None): + if attempt is None: + attempt = self.next_attempt(source_fingerprint) + now = datetime.now(timezone.utc).isoformat() + with self._connect() as connection: + cursor = connection.execute( + "INSERT INTO jobs(source_fingerprint, attempt_number, status, started_at) " + "VALUES (?, ?, ?, ?)", + (source_fingerprint, attempt, "pending", now), + ) + return cursor.lastrowid + + def next_attempt(self, source_fingerprint): + with self._connect() as connection: + row = connection.execute( + "SELECT COALESCE(MAX(attempt_number), 0) + 1 AS next_attempt " + "FROM jobs WHERE source_fingerprint=?", + (source_fingerprint,), + ).fetchone() + return int(row["next_attempt"]) + + def mark_authorization_started(self, job_id): + with self._connect() as connection: + connection.execute( + "UPDATE jobs SET authorization_started=1 " + "WHERE job_id=? AND status='pending'", + (job_id,), + ) + + def finish(self, job_id, status, reason_code, sink_receipt_fingerprint=None): + status_value = status.value if isinstance(status, JobStatus) else str(status) + now = datetime.now(timezone.utc).isoformat() + with self._connect() as connection: + cursor = connection.execute( + "UPDATE jobs SET status=?, finished_at=?, reason_code=?, " + "sink_receipt_fingerprint=? " + "WHERE job_id=? AND status='pending'", + (status_value, now, reason_code, sink_receipt_fingerprint, job_id), + ) + if ( + cursor.rowcount == 1 + and status_value == JobStatus.IMPORTED.value + and sink_receipt_fingerprint + ): + connection.execute( + "INSERT OR IGNORE INTO credential_inventory(" + "sink_receipt_fingerprint, state) VALUES (?, 'available')", + (sink_receipt_fingerprint,), + ) + + def claim_available(self, limit, batch_id): + limit = int(limit) + if limit <= 0: + return [] + batch_id = str(batch_id).strip() + if not batch_id: + raise ValueError("batch_id must not be empty") + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + rows = connection.execute( + """ + WITH latest_import AS ( + SELECT sink_receipt_fingerprint, MAX(finished_at) AS finished_at + FROM jobs + WHERE status=? + AND COALESCE(TRIM(sink_receipt_fingerprint), '') <> '' + GROUP BY sink_receipt_fingerprint + ) + SELECT inventory.sink_receipt_fingerprint + FROM credential_inventory AS inventory + JOIN latest_import USING (sink_receipt_fingerprint) + WHERE inventory.state='available' + ORDER BY latest_import.finished_at DESC, + inventory.sink_receipt_fingerprint ASC + LIMIT ? + """, + (JobStatus.IMPORTED.value, limit), + ).fetchall() + fingerprints = [row["sink_receipt_fingerprint"] for row in rows] + connection.executemany( + "UPDATE credential_inventory SET state='claiming', batch_id=?, " + "claimed_at='', note='' " + "WHERE sink_receipt_fingerprint=? AND state='available'", + ((batch_id, fingerprint) for fingerprint in fingerprints), + ) + return fingerprints + + def mark_claimed(self, batch_id, note=""): + batch_id = str(batch_id).strip() + if not batch_id: + raise ValueError("batch_id must not be empty") + with self._connect() as connection: + cursor = connection.execute( + "UPDATE credential_inventory SET state='claimed', claimed_at=?, note=? " + "WHERE state='claiming' AND batch_id=?", + (datetime.now(timezone.utc).isoformat(), str(note), batch_id), + ) + return cursor.rowcount + + def inventory_counts(self): + counts = {"available": 0, "claiming": 0, "claimed": 0} + with self._connect() as connection: + rows = connection.execute( + "SELECT state, COUNT(*) AS count FROM credential_inventory GROUP BY state" + ) + for row in rows: + counts[row["state"]] = int(row["count"]) + return counts + + def pending_claims(self, batch_id=None): + query = ( + "SELECT sink_receipt_fingerprint, batch_id, note " + "FROM credential_inventory WHERE state='claiming'" + ) + params = () + if batch_id is not None: + query += " AND batch_id=?" + params = (str(batch_id),) + query += " ORDER BY batch_id, sink_receipt_fingerprint" + with self._connect() as connection: + rows = connection.execute(query, params) + return [dict(row) for row in rows] + + def recover_claiming(self, batch_id=None, *, note=""): + query = ( + "UPDATE credential_inventory SET state='available', claimed_at='', " + "batch_id='', note=? WHERE state='claiming'" + ) + params = [str(note)] + if batch_id is not None: + query += " AND batch_id=?" + params.append(str(batch_id)) + with self._connect() as connection: + cursor = connection.execute(query, params) + return cursor.rowcount + + def recover_pending(self, *, reason="recovered_pending"): + with self._connect() as connection: + connection.execute( + "UPDATE jobs SET status=?, finished_at=?, reason_code=? WHERE status='pending'", + (JobStatus.CANCELLED.value, datetime.now(timezone.utc).isoformat(), reason), + ) + + def has_imported(self, source_id): + with self._connect() as connection: + row = connection.execute( + "SELECT 1 FROM jobs WHERE source_fingerprint=? AND status=? LIMIT 1", + (self.fingerprint(source_id), JobStatus.IMPORTED.value), + ).fetchone() + return row is not None + + def imported_fingerprints(self): + with self._connect() as connection: + rows = connection.execute( + "SELECT DISTINCT source_fingerprint FROM jobs WHERE status=?", + (JobStatus.IMPORTED.value,), + ) + return {row["source_fingerprint"] for row in rows} + + def aggregate_counts(self): + with self._connect() as connection: + row = connection.execute( + """ + SELECT + COUNT(DISTINCT CASE WHEN authorization_started=1 OR status=? + THEN source_fingerprint END) AS attempted_unique, + COUNT(DISTINCT CASE WHEN status=? + THEN source_fingerprint END) AS imported_unique, + COUNT(CASE WHEN finished_at IS NOT NULL THEN 1 END) AS finalized_attempts, + COUNT(CASE WHEN status=? THEN 1 END) AS imported_attempts, + COUNT(CASE WHEN reason_code='rate_limited' THEN 1 END) AS rate_limited + FROM jobs + """, + ( + JobStatus.IMPORTED.value, + JobStatus.IMPORTED.value, + JobStatus.IMPORTED.value, + ), + ).fetchone() + return {key: int(row[key] or 0) for key in row.keys()} + + def get(self, job_id): + with self._connect() as connection: + row = connection.execute( + "SELECT source_fingerprint, attempt_number, status, started_at, finished_at, " + "reason_code, sink_receipt_fingerprint FROM jobs WHERE job_id=?", + (job_id,), + ).fetchone() + return dict(row) if row else None diff --git a/xai_enroller/models.py b/xai_enroller/models.py new file mode 100644 index 0000000..9e1ab4b --- /dev/null +++ b/xai_enroller/models.py @@ -0,0 +1,128 @@ +from dataclasses import dataclass +from enum import Enum +from typing import Optional + + +@dataclass(frozen=True) +class SourceRecord: + source_id: str + sso_token: str + cookies: tuple[dict, ...] = () + + def __repr__(self) -> str: + return "SourceRecord()" + + +@dataclass(frozen=True) +class DeviceFlow: + device_code: str + user_code: str + verification_url: str + expires_in: int + interval: float + token_endpoint: str = "https://auth.x.ai/oauth2/token" + + def __repr__(self) -> str: + return "DeviceFlow()" + + +@dataclass(frozen=True) +class OAuthCredential: + access_token: str + refresh_token: str + id_token: Optional[str] + token_type: Optional[str] + expires_in: int + expires_at: str + last_refresh: str + subject: Optional[str] + token_endpoint: str + + def __repr__(self) -> str: + return "OAuthCredential()" + + +@dataclass(frozen=True) +class SinkReceipt: + fingerprint: str + + def __repr__(self) -> str: + return "SinkReceipt()" + + +class AuthorizationStatus(str, Enum): + AUTHORIZED = "authorized" + NEEDS_BROWSER = "needs_browser" + NEEDS_INTERACTION = "needs_interaction" + + +@dataclass(frozen=True) +class AuthorizationResult: + status: AuthorizationStatus + reason_code: str + + +class JobStatus(str, Enum): + IMPORTED = "imported" + NEEDS_BROWSER = "needs_browser" + NEEDS_INTERACTION = "needs_interaction" + SOURCE_INVALID = "source_invalid" + OAUTH_DENIED = "oauth_denied" + OAUTH_EXPIRED = "oauth_expired" + OAUTH_REJECTED = "oauth_rejected" + SINK_FAILED = "sink_failed" + TRANSPORT_FAILED = "transport_failed" + TIMEOUT = "timeout" + CANCELLED = "cancelled" + + +@dataclass(frozen=True) +class JobResult: + source_id: str + status: JobStatus + reason_code: str + attempt_number: int = 1 + sink_receipt_fingerprint: Optional[str] = None + + def __repr__(self) -> str: + return ( + f"JobResult(source_id=, status={self.status.value!r}, " + f"reason_code={self.reason_code!r}, attempt_number={self.attempt_number})" + ) + + +class PipelineState(str, Enum): + QUEUED = "queued" + PREPARED = "prepared" + ACTIVE = "active" + RETRY_WAITING = "retry_waiting" + IMPORTED = "imported" + TERMINAL = "terminal" + + +@dataclass(frozen=True) +class PreparedJob: + source: SourceRecord + source_fingerprint: str + flow: DeviceFlow + flow_created_monotonic: float + job_id: int + attempt_number: int + task_number: Optional[int] = None + + def __repr__(self) -> str: + return ( + "PreparedJob(source=, flow=, " + f"attempt_number={self.attempt_number})" + ) + + +@dataclass(frozen=True) +class CompletionJob: + prepared: PreparedJob + + def __repr__(self) -> str: + return ( + "CompletionJob(source=, flow=, " + f"attempt_number={self.prepared.attempt_number})" + ) diff --git a/xai_enroller/protocol.py b/xai_enroller/protocol.py new file mode 100644 index 0000000..f6f6bc4 --- /dev/null +++ b/xai_enroller/protocol.py @@ -0,0 +1,181 @@ +import asyncio +import base64 +import binascii +import json +import time +from dataclasses import dataclass +from datetime import datetime, timezone +from typing import Awaitable, Callable +from urllib.parse import urlencode + +import httpx + +from .models import DeviceFlow, OAuthCredential + + +@dataclass(frozen=True) +class XAIProfile: + client_id: str + scope: str + version: str = "xai-device-v1" + + @classmethod + def default(cls): + return cls( + client_id="b1a00492-073a-47ea-816f-4c329264a828", + scope="openid profile email offline_access grok-cli:access api:access", + version="grok-cli-device-v1", + ) + + +class XAIProtocol: + DISCOVERY_URL = "https://auth.x.ai/.well-known/openid-configuration" + MAX_BODY = 64 * 1024 + + def __init__( + self, + client: httpx.AsyncClient, + profile: XAIProfile, + sleep=None, + default_poll_interval=5.0, + ): + self.client = client + self.profile = profile + self.sleep = sleep or asyncio.sleep + self.default_poll_interval = float(default_poll_interval) + self._discovered = None + + @staticmethod + def _allowed_url(value): + from urllib.parse import urlparse + + parsed = urlparse(value) + hostname = parsed.hostname + return bool(hostname) and parsed.scheme == "https" and ( + hostname == "x.ai" or hostname.endswith(".x.ai") + ) + + async def _json(self, response): + body = await response.aread() + if len(body) > self.MAX_BODY: + raise RuntimeError("response body exceeds limit") + try: + return json.loads(body) + except (TypeError, ValueError) as exc: + raise RuntimeError("invalid JSON response") from exc + + async def discover(self): + response = await self.client.get(self.DISCOVERY_URL, follow_redirects=False) + if response.status_code // 100 != 2: + raise RuntimeError("discovery rejected") + document = await self._json(response) + device_endpoint = document.get("device_authorization_endpoint") + token_endpoint = document.get("token_endpoint") + if not device_endpoint or not token_endpoint or not all( + self._allowed_url(url) for url in (device_endpoint, token_endpoint) + ): + raise ValueError("discovery endpoint failed x.ai HTTPS allowlist") + self._discovered = (device_endpoint, token_endpoint) + return self._discovered + + async def start_device_flow(self): + device_endpoint, token_endpoint = await self.discover() + response = await self.client.post( + device_endpoint, + data={"client_id": self.profile.client_id, "scope": self.profile.scope}, + follow_redirects=False, + ) + if response.status_code // 100 != 2: + raise RuntimeError("device authorization rejected") + document = await self._json(response) + try: + device_code = document["device_code"] + user_code = document["user_code"] + base_url = document.get("verification_uri") or document["verification_url"] + expires_in = int(document["expires_in"]) + except (KeyError, TypeError, ValueError) as exc: + raise RuntimeError("invalid device authorization response") from exc + if not self._allowed_url(base_url): + raise ValueError("verification URL failed x.ai HTTPS allowlist") + interval = float(document.get("interval", self.default_poll_interval)) + verification_url = document.get("verification_uri_complete") + if verification_url is None: + verification_url = ( + f"{base_url}{'&' if '?' in base_url else '?'}" + f"{urlencode({'user_code': user_code})}" + ) + if not self._allowed_url(verification_url): + raise ValueError("verification URL failed x.ai HTTPS allowlist") + return DeviceFlow(device_code, user_code, verification_url, expires_in, interval, token_endpoint) + + async def poll_token(self, *, endpoint, flow, timeout): + deadline = time.monotonic() + timeout + interval = max(0.0, float(flow.interval)) + while time.monotonic() < deadline: + response = await self.client.post( + endpoint, + data={ + "client_id": self.profile.client_id, + "device_code": flow.device_code, + "grant_type": "urn:ietf:params:oauth:grant-type:device_code", + }, + follow_redirects=False, + ) + document = await self._json(response) + if response.status_code // 100 == 2: + return self._credential(document, endpoint) + error = document.get("error") + if error in {"authorization_pending", "slow_down"}: + interval += 1 if error == "slow_down" else 0 + await self.sleep(interval) + continue + if error == "access_denied": + raise RuntimeError("oauth_denied") + if error == "expired_token": + raise RuntimeError("oauth_expired") + raise RuntimeError("oauth_rejected") + raise RuntimeError("oauth_expired") + + @staticmethod + def _jwt_subject(token): + if not isinstance(token, str): + return None + parts = token.split(".") + if len(parts) != 3 or not parts[1]: + return None + payload = parts[1] + "=" * (-len(parts[1]) % 4) + try: + claims = json.loads(base64.urlsafe_b64decode(payload).decode("utf-8")) + except (binascii.Error, UnicodeDecodeError, ValueError): + return None + if not isinstance(claims, dict): + return None + for claim in ("sub", "principal_id"): + value = claims.get(claim) + if isinstance(value, str) and value: + return value + return None + + @staticmethod + def _credential(document, endpoint): + if not document.get("access_token") or not document.get("refresh_token"): + raise RuntimeError("oauth_rejected") + now = datetime.now(timezone.utc) + expires_in = int(document.get("expires_in", 3600)) + expires_at = datetime.fromtimestamp(now.timestamp() + expires_in, timezone.utc).isoformat().replace("+00:00", "Z") + subject = ( + XAIProtocol._jwt_subject(document.get("id_token")) + or XAIProtocol._jwt_subject(document["access_token"]) + or document.get("sub") + ) + return OAuthCredential( + access_token=document["access_token"], + refresh_token=document["refresh_token"], + id_token=document.get("id_token"), + token_type=document.get("token_type"), + expires_in=expires_in, + expires_at=expires_at, + last_refresh=now.isoformat().replace("+00:00", "Z"), + subject=subject, + token_endpoint=endpoint, + ) diff --git a/xai_enroller/remote_stream.py b/xai_enroller/remote_stream.py new file mode 100644 index 0000000..4c5bb99 --- /dev/null +++ b/xai_enroller/remote_stream.py @@ -0,0 +1,630 @@ +"""Atomic local snapshots of password-free registration sessions.""" + +import asyncio +import json +import os +import shlex +import sys +import tempfile +from contextlib import suppress +from pathlib import Path + +from .models import SourceRecord + +MAX_SESSION_RECORD_BYTES = 256 * 1024 + + +class RemoteStreamError(RuntimeError): + """A classified stream failure whose text never includes remote output.""" + + +class SSHSnapshotSynchronizer: + """Atomically refresh one validated, password-free local JSONL snapshot.""" + + MAX_STDERR_BYTES = 16 * 1024 + + def __init__( + self, + host, + destination, + *, + remote_root="/opt/grok-free-register", + identity_file=None, + process_factory=asyncio.create_subprocess_exec, + fingerprint=None, + ): + self.host = host + self.destination = Path(destination) + self.remote_root = remote_root + self.identity_file = identity_file + self.process_factory = process_factory + self.fingerprint = fingerprint or (lambda source_id: source_id) + self.snapshot_fingerprints = None + self._process = None + + def _input_generation(self): + return None + + def _input_generation_unchanged(self, generation): + return True + + def _empty_snapshot_allowed(self): + return False + + def _command(self): + return ( + f"cd {shlex.quote(self.remote_root)} && " + "python3 scripts/export_registered_sessions.py " + "keys/auth-sessions.jsonl keys/accounts.txt" + ) + + def _args(self): + args = [ + "ssh", + "-C", + "-T", + "-o", + "BatchMode=yes", + "-o", + "ConnectTimeout=15", + "-o", + "ServerAliveInterval=15", + "-o", + "ServerAliveCountMax=3", + ] + if self.identity_file: + args.extend(["-i", self.identity_file]) + args.extend(["--", self.host, self._command()]) + return args + + async def _read_stderr(self, stream): + retained = 0 + while True: + chunk = await stream.read(4096) + if not chunk: + return + retained = min(self.MAX_STDERR_BYTES, retained + len(chunk)) + + async def _terminate(self, process): + if process is None or process.returncode is not None: + return + with suppress(ProcessLookupError): + process.terminate() + try: + await asyncio.wait_for(process.wait(), timeout=3) + except TimeoutError: + with suppress(ProcessLookupError): + process.kill() + await process.wait() + + async def close(self): + await self._terminate(self._process) + + async def sync_once(self): + self.destination.parent.mkdir(mode=0o700, parents=True, exist_ok=True) + os.chmod(self.destination.parent, 0o700) + fd, temporary_name = tempfile.mkstemp( + prefix=f".{self.destination.name}.", + suffix=".tmp", + dir=self.destination.parent, + ) + process = None + stderr_task = None + try: + os.fchmod(fd, 0o600) + input_generation = self._input_generation() + process = await self.process_factory( + *self._args(), + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + limit=MAX_SESSION_RECORD_BYTES + 1, + ) + self._process = process + stderr_task = asyncio.create_task(self._read_stderr(process.stderr)) + with os.fdopen(fd, "wb") as stream: + fd = -1 + record_count = 0 + snapshot_fingerprints = set() + while True: + try: + raw = await process.stdout.readline() + except (ValueError, asyncio.LimitOverrunError) as exc: + raise ValueError("invalid remote session snapshot") from exc + if not raw: + break + if ( + not raw.endswith(b"\n") + or len(raw) - 1 > MAX_SESSION_RECORD_BYTES + ): + raise ValueError("invalid remote session snapshot") + record = parse_session_document(raw[:-1]) + snapshot_fingerprints.add(self.fingerprint(record.source_id)) + stream.write(raw) + record_count += 1 + returncode = await process.wait() + if stderr_task is not None: + await stderr_task + stderr_task = None + if returncode != 0: + raise RemoteStreamError("remote snapshot export failed") + if record_count == 0 and not self._empty_snapshot_allowed(): + raise RemoteStreamError("remote snapshot export was empty") + if not self._input_generation_unchanged(input_generation): + raise RemoteStreamError("snapshot inputs changed during export") + stream.flush() + os.fsync(stream.fileno()) + os.replace(temporary_name, self.destination) + os.chmod(self.destination, 0o600) + directory_fd = os.open(self.destination.parent, os.O_RDONLY) + try: + os.fsync(directory_fd) + finally: + os.close(directory_fd) + self.snapshot_fingerprints = frozenset(snapshot_fingerprints) + return True + except asyncio.CancelledError: + raise + except Exception: + return False + finally: + if fd >= 0: + os.close(fd) + if stderr_task is not None: + if not stderr_task.done(): + stderr_task.cancel() + await asyncio.gather(stderr_task, return_exceptions=True) + with suppress(Exception): + await self._terminate(process) + if self._process is process: + self._process = None + with suppress(FileNotFoundError): + os.unlink(temporary_name) + + +class LocalSnapshotSynchronizer(SSHSnapshotSynchronizer): + """Export registration data from one local project into an atomic snapshot.""" + + def __init__( + self, + register_root, + destination, + *, + process_factory=asyncio.create_subprocess_exec, + fingerprint=None, + python_executable=sys.executable, + exporter_path=None, + ): + super().__init__( + "", + destination, + process_factory=process_factory, + fingerprint=fingerprint, + ) + self.register_root = Path(register_root) + self.python_executable = str(python_executable) + self.exporter_path = Path(exporter_path or ( + Path(__file__).resolve().parents[1] + / "scripts" + / "export_registered_sessions.py" + )) + self.sessions_path = self.register_root / "keys" / "auth-sessions.jsonl" + self.accounts_path = self.register_root / "keys" / "accounts.txt" + + @staticmethod + def _path_generation(path): + try: + stat_result = path.stat() + except FileNotFoundError: + return None + return ( + stat_result.st_dev, + stat_result.st_ino, + stat_result.st_size, + stat_result.st_mtime_ns, + ) + + def _input_generation(self): + return ( + self._path_generation(self.sessions_path), + self._path_generation(self.accounts_path), + ) + + def _input_generation_unchanged(self, generation): + return self._input_generation() == generation + + def _empty_snapshot_allowed(self): + return True + + def _args(self): + return [ + self.python_executable, + str(self.exporter_path), + str(self.sessions_path), + str(self.accounts_path), + ] + + +class DiskSnapshotSource: + """Consume immutable local snapshot generations under queue backpressure.""" + + def __init__( + self, + path, + *, + synchronizer=None, + sync_seconds=30.0, + poll_seconds=0.25, + sleep=asyncio.sleep, + event_callback=None, + fingerprint=None, + ): + self.path = Path(path) + self.synchronizer = synchronizer + self.sync_seconds = float(sync_seconds) + self.poll_seconds = float(poll_seconds) + self.sleep = sleep + self.event_callback = event_callback + self.fingerprint = fingerprint or (lambda source_id: source_id) + self._snapshot_fingerprints = None + self._last_reported_snapshot_fingerprints = None + self._closed = False + self._sync_task = None + self._last_sync_ok = None + + def _emit(self, kind, data): + if self.event_callback is not None: + with suppress(Exception): + self.event_callback(kind, data) + + @property + def snapshot_fingerprints(self): + if self.synchronizer is not None: + return self.synchronizer.snapshot_fingerprints + return self._snapshot_fingerprints + + async def _sync_loop(self): + while not self._closed: + refreshed = await self.synchronizer.sync_once() + current = self.synchronizer.snapshot_fingerprints + if refreshed != self._last_sync_ok: + self._emit( + "source_connected" if refreshed else "source_disconnected", + {} if refreshed else {"reason": "snapshot_sync_failed"}, + ) + self._last_sync_ok = refreshed + if refreshed and current is not None: + previous = self._last_reported_snapshot_fingerprints + if previous is not None: + added = current - previous + if added: + self._emit( + "source_updated", + {"new": len(added), "total": len(current)}, + ) + self._last_reported_snapshot_fingerprints = current + try: + await asyncio.wait_for(self._wait_closed(), timeout=self.sync_seconds) + except TimeoutError: + pass + + async def _wait_closed(self): + while not self._closed: + await self.sleep(min(0.25, self.sync_seconds)) + + async def close(self): + self._closed = True + if self._sync_task is not None: + self._sync_task.cancel() + with suppress(asyncio.CancelledError): + await self._sync_task + self._sync_task = None + if self.synchronizer is not None: + await self.synchronizer.close() + + @staticmethod + def _generation(stat_result): + return ( + stat_result.st_dev, + stat_result.st_ino, + stat_result.st_mtime_ns, + stat_result.st_size, + ) + + async def records(self): + if self.synchronizer is not None and self._sync_task is None: + self._sync_task = asyncio.create_task(self._sync_loop()) + consumed_generation = None + while not self._closed: + try: + current = self.path.stat() + except FileNotFoundError: + await self.sleep(self.poll_seconds) + continue + generation = self._generation(current) + if generation == consumed_generation: + await self.sleep(self.poll_seconds) + continue + try: + stream = self.path.open("rb") + except FileNotFoundError: + continue + with stream: + opened_generation = self._generation(os.fstat(stream.fileno())) + generation_fingerprints = set() + valid_generation = True + while not self._closed: + raw = await asyncio.to_thread( + stream.readline, MAX_SESSION_RECORD_BYTES + 2 + ) + if not raw: + break + if ( + not raw.endswith(b"\n") + or len(raw) - 1 > MAX_SESSION_RECORD_BYTES + ): + self._emit( + "source_record_rejected", {"reason": "invalid_record"} + ) + valid_generation = False + break + try: + record = parse_session_document(raw[:-1]) + except ValueError: + self._emit( + "source_record_rejected", {"reason": "invalid_record"} + ) + valid_generation = False + continue + generation_fingerprints.add(self.fingerprint(record.source_id)) + yield record + if valid_generation: + self._snapshot_fingerprints = frozenset(generation_fingerprints) + consumed_generation = opened_generation + + +def parse_session_document(raw: bytes | str) -> SourceRecord: + if len(raw) > MAX_SESSION_RECORD_BYTES: + raise ValueError("invalid remote session record") + try: + document = json.loads(raw) + source_id = document["email"] + raw_cookies = document["cookies"] + except (UnicodeDecodeError, TypeError, ValueError, KeyError) as exc: + raise ValueError("invalid remote session record") from exc + if not isinstance(source_id, str) or not source_id: + raise ValueError("invalid remote session record") + try: + source_id.encode("utf-8") + except UnicodeEncodeError as exc: + raise ValueError("invalid remote session record") from exc + if not isinstance(raw_cookies, list) or not raw_cookies: + raise ValueError("invalid remote session record") + + cookies = [] + sso_token = "" + fallback_sso_token = "" + allowed = { + "name", + "value", + "url", + "domain", + "path", + "expires", + "httpOnly", + "secure", + "sameSite", + } + for raw_cookie in raw_cookies: + if not isinstance(raw_cookie, dict): + raise ValueError("invalid remote session record") + cookie = {key: raw_cookie[key] for key in allowed if key in raw_cookie} + name = cookie.get("name") + value = cookie.get("value") + if not isinstance(name, str) or not name or not isinstance(value, str) or not value: + raise ValueError("invalid remote session record") + scope = cookie.get("url") or cookie.get("domain") + if not isinstance(scope, str) or not scope: + raise ValueError("invalid remote session record") + try: + name.encode("utf-8") + value.encode("utf-8") + scope.encode("utf-8") + except UnicodeEncodeError as exc: + raise ValueError("invalid remote session record") from exc + if name == "sso" and not sso_token: + sso_token = value + elif name == "sso-rw" and not fallback_sso_token: + fallback_sso_token = value + cookies.append(cookie) + sso_token = sso_token or fallback_sso_token + if not sso_token: + raise ValueError("invalid remote session record") + return SourceRecord(source_id, sso_token, tuple(cookies)) + + +class RemoteSessionStream: + """Yield full snapshots and appends from one reconnecting SSH child.""" + + MAX_STDERR_BYTES = 16 * 1024 + MAX_RECORD_BYTES = MAX_SESSION_RECORD_BYTES + RECONNECT_DELAYS = (1.0, 2.0, 5.0, 10.0, 30.0) + + def __init__( + self, + host: str, + *, + remote_root: str = "/opt/grok-free-register", + identity_file: str | None = None, + process_factory=asyncio.create_subprocess_exec, + sleep=asyncio.sleep, + event_callback=None, + ): + self.host = host + self.remote_root = remote_root + self.identity_file = identity_file + self.process_factory = process_factory + self.sleep = sleep + self.event_callback = event_callback + self._process = None + self._closed = False + self._last_disconnect_reason = None + + def _emit(self, kind, data): + if self.event_callback is None: + return + try: + self.event_callback(kind, data) + except Exception: + pass + + def _command(self): + return ( + f"cd {shlex.quote(self.remote_root)} && " + "python3 -u scripts/export_registered_sessions.py --follow " + "keys/auth-sessions.jsonl keys/accounts.txt" + ) + + def _args(self): + args = [ + "ssh", + "-C", + "-T", + "-o", + "BatchMode=yes", + "-o", + "ConnectTimeout=15", + "-o", + "ServerAliveInterval=15", + "-o", + "ServerAliveCountMax=3", + ] + if self.identity_file: + args.extend(["-i", self.identity_file]) + args.extend(["--", self.host, self._command()]) + return args + + async def _read_stderr(self, stream): + retained = bytearray() + while True: + chunk = await stream.read(4096) + if not chunk: + break + remaining = self.MAX_STDERR_BYTES - len(retained) + if remaining > 0: + retained.extend(chunk[:remaining]) + return bytes(retained) + + @staticmethod + def _classify_disconnect(returncode, stderr): + normalized = stderr.lower() + if b"permission denied" in normalized or b"host key verification failed" in normalized: + return "ssh_auth_failed" + if b"could not resolve" in normalized or b"name or service not known" in normalized: + return "ssh_resolution_failed" + if any( + marker in normalized + for marker in (b"connection refused", b"connection timed out", b"no route to host") + ): + return "ssh_connection_failed" + if returncode == 3: + return "remote_snapshot_changed" + if returncode == 4: + return "remote_data_invalid" + return "remote_stream_closed" + + async def _terminate_process(self, process): + if process is None or process.returncode is not None: + return + with suppress(ProcessLookupError): + process.terminate() + try: + await asyncio.wait_for(process.wait(), timeout=3) + except TimeoutError: + with suppress(ProcessLookupError): + process.kill() + await process.wait() + + async def close(self): + self._closed = True + await self._terminate_process(self._process) + + async def records(self): + reconnect_index = 0 + while not self._closed: + process = None + stderr_task = None + yielded = False + reason = "remote_stream_closed" + try: + process = await self.process_factory( + *self._args(), + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + limit=self.MAX_RECORD_BYTES + 1, + ) + self._process = process + stderr_task = asyncio.create_task(self._read_stderr(process.stderr)) + self._emit("source_connected", {}) + while not self._closed: + try: + raw = await process.stdout.readline() + except (ValueError, asyncio.LimitOverrunError): + reason = "remote_record_too_large" + break + if not raw: + break + if ( + not raw.endswith(b"\n") + or len(raw) - 1 > self.MAX_RECORD_BYTES + ): + reason = "remote_record_too_large" + break + try: + record = parse_session_document(raw[:-1]) + except ValueError: + self._emit("source_record_rejected", {"reason": "invalid_record"}) + continue + if not yielded: + self._last_disconnect_reason = None + yielded = True + reconnect_index = 0 + yield record + + # A follow exporter is intentionally long-lived. Once framing is + # invalid, waiting for it to exit on its own would hang this source + # forever, so terminate it before collecting the exit status. + if reason != "remote_stream_closed": + await self._terminate_process(process) + if process.returncode is None: + await process.wait() + stderr = await stderr_task + if reason == "remote_stream_closed": + reason = self._classify_disconnect(process.returncode, stderr) + except asyncio.CancelledError: + raise + except (OSError, RuntimeError): + reason = "ssh_start_failed" + finally: + if stderr_task is not None and not stderr_task.done(): + stderr_task.cancel() + with suppress(asyncio.CancelledError): + await stderr_task + await self._terminate_process(process) + if self._process is process: + self._process = None + + if self._closed: + break + if reason != self._last_disconnect_reason: + self._emit("source_disconnected", {"reason": reason}) + self._last_disconnect_reason = reason + if reason == "remote_snapshot_changed": + delay = 0.1 + else: + delay = self.RECONNECT_DELAYS[min(reconnect_index, len(self.RECONNECT_DELAYS) - 1)] + if not yielded: + reconnect_index += 1 + await self.sleep(delay) + + +PersistentSSHSource = RemoteSessionStream diff --git a/xai_enroller/service.py b/xai_enroller/service.py new file mode 100644 index 0000000..f4f715b --- /dev/null +++ b/xai_enroller/service.py @@ -0,0 +1,883 @@ +"""Interactive composition for the asynchronous local authentication service.""" + +import asyncio +import os +import secrets +import sys +import tempfile +from contextlib import suppress +from dataclasses import dataclass +from pathlib import Path + +from dotenv import load_dotenv + +from .inventory import InventoryError + + +DEFAULT_LOCAL_AUTH_DIR = Path.home() / "Downloads" / "grok-free-register-auth" +AUTHENTICATED_DIRNAME = "authenticated" +CLAIMED_DIRNAME = "claimed" + + +class AuthServiceConfigurationError(ValueError): + pass + + +class InteractiveCommandPrompt: + """Stdlib-only line editor that keeps commands visible during events.""" + + CLEAR_LINE = "\r\x1b[2K" + + def __init__( + self, + *, + input_stream=None, + output=None, + interactive=None, + prompt="认证> ", + ): + self.input_stream = input_stream or sys.stdin + self.output = output or sys.stdout + self.prompt = prompt + if interactive is None: + interactive = bool( + getattr(self.input_stream, "isatty", lambda: False)() + and getattr(self.output, "isatty", lambda: False)() + ) + self.interactive = interactive + self.buffer = "" + self.active = False + self._callback = None + self._fd = None + self._original_terminal = None + + def _write(self, text): + try: + self.output.write(text) + self.output.flush() + except (OSError, ValueError): + pass + + def _redraw(self): + if self.active and self.interactive: + self._write(f"{self.CLEAR_LINE}{self.prompt}{self.buffer}") + + def start(self, callback): + self._callback = callback + self.active = True + self._redraw() + + def feed(self, text): + for character in text: + if character in {"\r", "\n"}: + command = self.buffer + self.buffer = "" + self._write("\n") + if command.strip() and self._callback is not None: + self._callback(command) + if self.active and self.interactive: + self._write(self.prompt) + elif character in {"\b", "\x7f"}: + self.buffer = self.buffer[:-1] + self._redraw() + elif character == "\x04": + if self._callback is not None: + self._callback(None) + elif character.isprintable(): + self.buffer += character + self._redraw() + + def write_event(self, message): + if self.active and self.interactive: + self._write(f"{self.CLEAR_LINE}{message}\n{self.prompt}{self.buffer}") + else: + self._write(f"{message}\n") + + def attach(self, loop, callback): + self.start(callback) + self._fd = self.input_stream.fileno() + if self.interactive: + import termios + import tty + + self._original_terminal = termios.tcgetattr(self._fd) + tty.setcbreak(self._fd) + + def stdin_ready(): + data = os.read(self._fd, 256) + if data: + self.feed(data.decode(errors="ignore")) + elif self._callback is not None: + loop.remove_reader(self._fd) + self._callback(None) + + else: + + def stdin_ready(): + line = self.input_stream.readline() + if line == "": + if self._callback is not None: + loop.remove_reader(self._fd) + self._callback(None) + return + if self._callback is not None: + self._callback(line) + + loop.add_reader(self._fd, stdin_ready) + + def close(self, loop): + if self._fd is not None: + with suppress(Exception): + loop.remove_reader(self._fd) + if self._original_terminal is not None and self._fd is not None: + import termios + + with suppress(Exception): + termios.tcsetattr( + self._fd, termios.TCSADRAIN, self._original_terminal + ) + if self.active and self.interactive: + self._write(f"{self.CLEAR_LINE}\n") + self.active = False + + +def resolve_auth_log_mode(argv=None, env=None): + argv = list(sys.argv[1:] if argv is None else argv) + env = dict(os.environ if env is None else env) + unknown = [argument for argument in argv if argument != "--debug"] + if unknown: + raise AuthServiceConfigurationError( + f"unsupported auth service argument: {unknown[0]}" + ) + mode = (env.get("XAI_AUTH_SERVICE_LOG_MODE") or "user").strip().lower() + if "--debug" in argv: + mode = "debug" + if mode not in {"user", "debug"}: + raise AuthServiceConfigurationError( + "XAI_AUTH_SERVICE_LOG_MODE must be user or debug" + ) + return mode + + +def prepare_local_service_environment(env=None): + """Create private local persistence defaults without storing secrets in the repo.""" + merged = dict(os.environ if env is None else env) + destination = Path( + merged.get("XAI_ENROLLER_LOCAL_AUTH_DIR", DEFAULT_LOCAL_AUTH_DIR) + ).expanduser() + destination.mkdir(mode=0o700, parents=True, exist_ok=True) + os.chmod(destination, 0o700) + merged["XAI_ENROLLER_SINK"] = "local" + merged["XAI_ENROLLER_LOCAL_AUTH_DIR"] = str(destination) + merged.setdefault( + "XAI_ENROLLER_LEDGER_PATH", str(destination / "enrollment-ledger.db") + ) + if not merged.get("XAI_ENROLLER_SOURCE_SALT"): + salt_file = destination / ".ledger-salt" + try: + salt = salt_file.read_text(encoding="utf-8").strip() + except FileNotFoundError: + salt = secrets.token_hex(32) + fd, temporary_name = tempfile.mkstemp( + prefix=".ledger-salt.", suffix=".tmp", dir=destination, text=True + ) + try: + os.fchmod(fd, 0o600) + with os.fdopen(fd, "w", encoding="utf-8") as stream: + stream.write(salt + "\n") + stream.flush() + os.fsync(stream.fileno()) + os.replace(temporary_name, salt_file) + finally: + if os.path.exists(temporary_name): + os.unlink(temporary_name) + os.chmod(salt_file, 0o600) + merged["XAI_ENROLLER_SOURCE_SALT"] = salt + return merged + + +@dataclass(frozen=True) +class AuthServiceSettings: + source_kind: str + ssh_host: str | None + register_root: str + remote_root: str + identity_file: str | None + sync_seconds: int = 30 + retry_seconds: int = 60 + min_authorization_interval_seconds: float = 10.0 + + @classmethod + def from_environ(cls, env=None): + env = dict(os.environ if env is None else env) + ssh_host = (env.get("XAI_AUTH_SERVICE_SSH_HOST") or "").strip() + requested_source = (env.get("XAI_AUTH_SERVICE_SOURCE") or "auto").strip().lower() + if requested_source not in {"auto", "local", "ssh"}: + raise AuthServiceConfigurationError( + "XAI_AUTH_SERVICE_SOURCE must be auto, local, or ssh" + ) + source_kind = ( + "ssh" if requested_source == "auto" and ssh_host else + "local" if requested_source == "auto" else + requested_source + ) + if source_kind == "ssh" and not ssh_host: + raise AuthServiceConfigurationError("XAI_AUTH_SERVICE_SSH_HOST is required") + register_root = ( + env.get("XAI_AUTH_SERVICE_REGISTER_ROOT") + or str(Path(__file__).resolve().parents[1]) + ).strip() + remote_root = ( + env.get("XAI_AUTH_SERVICE_REMOTE_ROOT") or "/opt/grok-free-register" + ).strip() + identity_file = (env.get("XAI_AUTH_SERVICE_SSH_IDENTITY") or "").strip() or None + try: + sync_seconds = int(env.get("XAI_AUTH_SERVICE_SYNC_SEC", "30")) + retry_seconds = int(env.get("XAI_AUTH_SERVICE_RETRY_SEC", "60")) + min_authorization_interval_seconds = float( + env.get("XAI_AUTH_SERVICE_MIN_INTERVAL_SEC", "10") + ) + except ValueError as exc: + raise AuthServiceConfigurationError( + "auth service intervals must be numeric" + ) from exc + if not 5 <= sync_seconds <= 3600: + raise AuthServiceConfigurationError( + "XAI_AUTH_SERVICE_SYNC_SEC must be between 5 and 3600" + ) + if not 30 <= retry_seconds <= 86400: + raise AuthServiceConfigurationError( + "XAI_AUTH_SERVICE_RETRY_SEC must be between 30 and 86400" + ) + if not 0 <= min_authorization_interval_seconds <= 3600: + raise AuthServiceConfigurationError( + "XAI_AUTH_SERVICE_MIN_INTERVAL_SEC must be between 0 and 3600" + ) + return cls( + source_kind=source_kind, + ssh_host=ssh_host or None, + register_root=register_root, + remote_root=remote_root, + identity_file=identity_file, + sync_seconds=sync_seconds, + retry_seconds=retry_seconds, + min_authorization_interval_seconds=min_authorization_interval_seconds, + ) + + +class EventTerminal: + """Mode-aware aggregate renderer; identifiers and credentials stay hidden.""" + + _SAFE_REASONS = { + "already_imported", + "authorization_transport_failed", + "browser_error", + "challenge_required", + "confirmation_timeout", + "device_flow_failed", + "device_flow_invalid", + "device_flow_refresh_failed", + "imported", + "internal_error", + "oauth_denied", + "oauth_expired", + "oauth_rejected", + "operator_cancelled", + "rate_limited", + "sink_failed", + "sink_unconfigured", + "snapshot_sync_failed", + "token_failed", + "token_transport_failed", + "unknown_page", + } + _TRANSIENT_REASONS = { + "authorization_transport_failed", + "browser_error", + "challenge_required", + "confirmation_timeout", + "device_flow_failed", + "device_flow_refresh_failed", + "sink_failed", + "token_failed", + "token_transport_failed", + "unknown_page", + } + + def __init__(self, mode="user", output=None): + if mode not in {"user", "debug"}: + raise ValueError("terminal mode must be user or debug") + self.mode = mode + self.output = output or self._print + + @staticmethod + def _print(message): + print(message, flush=True) + + @staticmethod + def _percentage(value): + return f"{100.0 * float(value):.1f}%" + + @staticmethod + def _rate(value): + return "—" if value is None else f"{float(value):.2f}/分" + + @staticmethod + def _pending(value): + return "—" if value is None else str(value) + + @staticmethod + def _task_suffix(value): + return "" if value is None else f" #{value}" + + def _format_user(self, kind, data): + if kind == "startup": + source = {"local": "本机", "ssh": "远端"}.get( + data.get("source_kind"), "等待同步" + ) + return ( + f"[✓] 本地认证服务已启动 | 来源 {source} | " + f"输出 {data['destination']} | 待处理 — | 可用 {data['available']}" + ) + if kind == "service_started": + return "[▶] 认证流水线运行中" + if kind == "service_stopped": + return "[■] 认证服务已停止" + if kind == "source_connected": + return "[✓] 账号来源已连接 | 本地快照可用" + if kind == "source_updated": + return f"[↻] 发现新账号 {data['new']} | 快照共 {data['total']}" + if kind == "source_disconnected": + return "[!] 账号来源暂时断开 | 继续使用上一份有效快照" + if kind == "source_record_rejected": + return "[!] 已忽略一条无效来源记录" + if kind == "rate_limited": + pace = data.get("min_authorization_interval_seconds") + pace_text = ( + f" | 恢复后间隔 {pace:.0f}秒" if pace is not None else "" + ) + return ( + f"[⏸] 触发限流 | {data['wait_seconds']}秒后单次探测{pace_text}" + ) + if kind == "rate_limit_cleared": + return f"[▶] 限流解除 | 实际等待 {data['elapsed_seconds']}秒" + if kind == "pacing_adjusted": + return ( + f"[•] 认证间隔调整为 {data['min_authorization_interval_seconds']:.0f}秒" + ) + if kind == "authorization_started": + return ( + f"[→] 开始认证 #{data['task_number']} | " + f"待处理 {self._pending(data.get('pending_total'))}" + ) + if kind == "result": + if data.get("status") == "imported": + return ( + f"[✓] 认证成功{self._task_suffix(data.get('task_number'))} | " + f"运行平均 {self._rate(data.get('lifetime_imports_per_minute'))} | " + f"累计 {data['imported_unique']} | 可用 {data.get('available', '—')}" + ) + reason = data.get("reason") + if reason == "rate_limited": + return None + action = "暂时失败,将自动重试" if reason in self._TRANSIENT_REASONS else "已跳过" + return f"[✗] 认证未完成{self._task_suffix(data.get('task_number'))} | {action}" + if kind == "control": + states = { + "paused": "[⏸] 认证服务已暂停", + "running": "[▶] 认证服务已恢复", + "cancelling": "[■] 正在取消当前任务", + "idle": "[•] 当前没有可取消的任务", + "stopping": "[■] 正在安全退出", + "usage: take ": "[!] 用法:take N,N 必须是正整数", + "inventory unavailable": "[!] 当前凭据库存不可用", + "help": ( + "[•] 命令 s 状态 | take N 取用 | p 暂停 | " + "r 恢复 | c 取消 | q 退出" + ), + "unknown": "[!] 未知命令 | 输入 help 查看命令", + } + return states.get(data.get("state"), "[•] 控制命令已处理") + if kind == "status": + state = {"running": "运行中", "paused": "已暂停", "stopping": "停止中"}.get( + data.get("state"), "未知" + ) + cooldown = ( + f"限流冷却 {data.get('cooldown_remaining_seconds', 0):.0f}秒" + if data.get("cooldown") + else "不限流" + ) + return ( + f"[•] 状态 {state} | 待处理 {self._pending(data.get('pending_total'))} | " + f"阶段 {data.get('active_stage', 'idle')} | " + f"运行平均 {self._rate(data.get('lifetime_imports_per_minute'))} | " + f"累计 {data.get('imported_unique', 0)} | " + f"可用 {data.get('available', 0)} | 已取用 {data.get('claimed', 0)} | {cooldown}" + ) + if kind == "inventory_taken": + return ( + f"[✓] 已取用 {data['moved']} 个凭据 | " + f"剩余可用 {data['available']} | " + f"位置 claimed/{data['batch_id']}/" + ) + if kind == "inventory_error": + return ( + f"[!] 凭据取用失败 | 可用 {data['available']} | " + f"处理中 {data['claiming']} | 已取用 {data['claimed']}" + ) + if kind == "pipeline_error": + return f"[!] 认证流水线异常 | 阶段 {data.get('stage', 'unknown')}" + return None + + def _format_debug(self, kind, data): + if kind == "service_started": + return ( + "• service: running; " + f"min_interval={data['min_authorization_interval_seconds']:.1f}s" + ) + if kind == "service_stopped": + return "• service: stopped" + if kind == "source_connected": + return "• source: local snapshot updated" + if kind == "source_updated": + return f"• source: new={data['new']}; total={data['total']}" + if kind == "source_disconnected": + return "⚠ source: snapshot_sync_failed; keeping previous snapshot" + if kind == "source_record_rejected": + return "⚠ source: invalid record rejected" + if kind == "rate_limited": + pace = data.get("min_authorization_interval_seconds") + pace_text = f"; resume_interval={pace:.1f}s" if pace is not None else "" + return ( + "⏸ authentication rate limited; " + f"next single probe in {data['wait_seconds']}s{pace_text}" + ) + if kind == "rate_limit_cleared": + return ( + "▶ authentication rate limit cleared after " + f"{data['elapsed_seconds']}s" + ) + if kind == "pacing_adjusted": + return ( + "• pacing: min_interval=" + f"{data['min_authorization_interval_seconds']:.1f}s" + f" ({data.get('reason', 'adjusted')})" + ) + if kind == "authorization_started": + return ( + "→ authentication: next task started; " + f"task={data['task_number']}; attempt={data['attempt_number']}; " + f"queued={data['source_queue']}; " + f"pending={self._pending(data.get('pending_total'))}" + ) + if kind == "result": + if data.get("status") == "imported": + return ( + "✓ authentication: imported; " + f"task={data.get('task_number')}; total={data['imported_unique']}; " + f"5m_rate={self._rate(data.get('five_minute_imports_per_minute'))}; " + f"attempt_success={self._percentage(data['attempt_success'])}; " + f"eventual_success={self._percentage(data['eventual_success'])}" + ) + reason = data.get("reason") + safe_reason = reason if reason in self._SAFE_REASONS else "unknown" + return ( + f"⚠ authentication: {data.get('status', 'unknown')} ({safe_reason}); " + f"task={data.get('task_number')}; attempt={data.get('attempt_number')}" + ) + if kind == "control": + state = data.get("state", "unknown") + safe_state = state if str(state).replace(" ", "").isalnum() else "unknown" + return f"• service: {safe_state}" + if kind == "status": + rate = data.get("five_minute_imports_per_minute") + rate_text = "unknown" if rate is None else f"{rate:.2f}/min" + lifetime_rate = data.get("lifetime_imports_per_minute") + lifetime_rate_text = ( + "unknown" + if lifetime_rate is None + else f"{lifetime_rate:.2f}/min" + ) + return ( + f"• status: {data['state']}; " + f"queues={data['source_queue']}/{data['prepared_queue']}/" + f"{data['completion_queue']}; active={data['active_stage']}; " + f"retry_waiting={data['retry_waiting']}; " + f"next_retry={data['next_retry_seconds']:.1f}s; " + f"started={data['authorization_starts']}; " + f"cooldown={str(data['cooldown']).lower()}; " + f"cooldown_remaining={data['cooldown_remaining_seconds']:.1f}s; " + f"probe={str(data['probe_in_flight']).lower()}; " + f"min_interval={data['min_authorization_interval_seconds']:.1f}s; " + f"pace_remaining={data['pacing_remaining_seconds']:.1f}s; " + f"imported={data['imported_unique']}; attempted={data['attempted_unique']}; " + f"rate_limited={data['rate_limited']}; 5m_rate={rate_text}; " + f"lifetime_rate={lifetime_rate_text}; " + f"available={data['available']}; claiming={data['claiming']}; " + f"claimed={data['claimed']}" + ) + if kind == "inventory_taken": + return ( + f"✓ inventory: claimed={data['moved']}; " + f"available={data['available']}; batch={data['batch_id']}" + ) + if kind == "inventory_error": + return ( + f"⚠ inventory: available={data['available']}; " + f"claiming={data['claiming']}; claimed={data['claimed']}" + ) + if kind == "pipeline_error": + reason = data.get("reason") + safe_reason = reason if reason in self._SAFE_REASONS else "unknown" + stage = data.get("stage", "unknown") + safe_stage = stage if str(stage).replace("_", "").isalnum() else "unknown" + return f"⚠ service: {safe_stage} stage stopped ({safe_reason})" + known = self._format_user(kind, data) + if known is not None: + return known + reason = data.get("reason") + suffix = f" reason={reason}" if reason in self._SAFE_REASONS else "" + safe_kind = kind if kind.replace("_", "").isalnum() else "unknown" + return f"• debug event={safe_kind}{suffix}" + + def emit(self, event): + kind, data = event + try: + message = ( + self._format_debug(kind, data) + if self.mode == "debug" + else self._format_user(kind, data) + ) + if message is not None: + self.output(message) + except Exception: + return + + +class AuthPipelineRunner: + """Map interactive controls onto the persistent pipeline lifecycle.""" + + def __init__(self, pipeline, emit, *, interval_seconds=None, inventory=None): + self.pipeline = pipeline + self.emit = emit + self.inventory = inventory + self.paused = False + self.current_cycle = None + + async def handle_command(self, command): + command = command.strip().lower() + if command == "p": + self.paused = True + self.pipeline.pause() + self.emit(("control", {"state": "paused"})) + elif command == "r": + self.paused = False + self.pipeline.resume() + self.emit(("control", {"state": "running"})) + elif command == "s": + status = self.pipeline.status() + status.update(self.pipeline.ledger.inventory_counts()) + self.emit(("status", status)) + elif command.startswith("take "): + parts = command.split() + if len(parts) != 2 or not parts[1].isdigit() or int(parts[1]) <= 0: + self.emit(("control", {"state": "usage: take "})) + return True + if self.inventory is None: + self.emit(("control", {"state": "inventory unavailable"})) + return True + try: + batch = await asyncio.to_thread(self.inventory.take, int(parts[1])) + except InventoryError as exc: + self.emit( + ( + "inventory_error", + { + "reason": str(exc), + **self.pipeline.ledger.inventory_counts(), + }, + ) + ) + return True + counts = self.pipeline.ledger.inventory_counts() + self.emit( + ( + "inventory_taken", + { + "batch_id": batch.batch_id, + "directory": str(batch.directory), + "moved": batch.moved, + **counts, + }, + ) + ) + elif command == "c": + cancelled = await self.pipeline.cancel_active() + self.emit( + ("control", {"state": "cancelling" if cancelled else "idle"}) + ) + elif command in {"help", "h", "?"}: + self.emit(("control", {"state": "help"})) + elif command in {"q", "quit", "exit"}: + self.pipeline.request_stop() + self.emit(("control", {"state": "stopping"})) + return False + elif command: + self.emit(("control", {"state": "unknown"})) + return True + + async def run(self): + self.current_cycle = asyncio.create_task(self.pipeline.run()) + try: + await self.current_cycle + finally: + self.current_cycle = None + + +async def _run_interactive(runner, terminal): + loop = asyncio.get_running_loop() + commands = asyncio.Queue() + prompt = InteractiveCommandPrompt() + previous_output = terminal.output + terminal.output = prompt.write_event + prompt.write_event( + "命令:s 状态 | take N 取用凭据 | p 暂停 | r 恢复 | c 取消当前任务 | q 退出" + ) + try: + prompt.attach(loop, commands.put_nowait) + except Exception: + prompt.close(loop) + terminal.output = previous_output + raise + worker = asyncio.create_task(runner.run()) + try: + while not worker.done(): + command_task = asyncio.create_task(commands.get()) + done, _pending = await asyncio.wait( + {worker, command_task}, return_when=asyncio.FIRST_COMPLETED + ) + if worker in done: + command_task.cancel() + with suppress(asyncio.CancelledError): + await command_task + await worker + break + command = command_task.result() + if command is None: + command = "q" + if not await runner.handle_command(command): + break + except asyncio.CancelledError: + runner.pipeline.request_stop() + raise + finally: + prompt.close(loop) + terminal.output = previous_output + runner.pipeline.request_stop() + await asyncio.gather(worker, return_exceptions=True) + + +async def main_async(*, log_mode="user"): + import httpx + + from .auth_pipeline import AuthPipeline + from .config import Settings + from .executors import PlaywrightExecutor + from .inventory import CredentialInventory + from .ledger import Ledger + from .protocol import XAIProfile, XAIProtocol + from .remote_stream import ( + DiskSnapshotSource, + LocalSnapshotSynchronizer, + SSHSnapshotSynchronizer, + ) + from .sinks import LocalAuthFileSink + + service_settings = AuthServiceSettings.from_environ() + merged = prepare_local_service_environment() + merged["XAI_ENROLLER_SOURCE_KIND"] = "remote" + merged["XAI_ENROLLER_AUTH_EXECUTOR"] = "playwright" + merged["XAI_ENROLLER_CONCURRENCY"] = "1" + settings = Settings.from_environ(merged) + terminal = EventTerminal(mode=log_mode) + + # CF prewarm via external FlareSolverr (:8191). Separate from REGISTER_PROXY. + try: + from grok_register.clearance import ( + format_prewarm_log, + httpx_proxy_mounts, + prewarm_clearance, + register_proxy_url, + ) + + prewarm = await asyncio.to_thread(prewarm_clearance) + terminal.emit(("control", {"state": format_prewarm_log(prewarm)})) + reg_proxy = register_proxy_url() + if reg_proxy: + terminal.emit(("control", {"state": f"register_proxy={reg_proxy}"})) + proxy_url = httpx_proxy_mounts() + except Exception as exc: # noqa: BLE001 + terminal.emit(("control", {"state": f"clearance_prewarm_skip: {exc}"})) + proxy_url = None + + # httpx honors REGISTER_PROXY/HTTP_PROXY only when passed explicitly + client = httpx.AsyncClient(proxy=proxy_url) if proxy_url else httpx.AsyncClient() + pipeline = None + try: + ledger = Ledger(settings.ledger_path, settings.source_salt) + snapshot_path = Path(settings.local_auth_dir) / "source-snapshot.jsonl" + if service_settings.source_kind == "ssh": + synchronizer = SSHSnapshotSynchronizer( + service_settings.ssh_host, + snapshot_path, + remote_root=service_settings.remote_root, + identity_file=service_settings.identity_file, + fingerprint=ledger.fingerprint, + ) + else: + synchronizer = LocalSnapshotSynchronizer( + service_settings.register_root, + snapshot_path, + fingerprint=ledger.fingerprint, + ) + source = DiskSnapshotSource( + snapshot_path, + synchronizer=synchronizer, + sync_seconds=service_settings.sync_seconds, + event_callback=lambda kind, data: terminal.emit((kind, data)), + fingerprint=ledger.fingerprint, + ) + protocol = XAIProtocol( + client, + XAIProfile.default(), + default_poll_interval=settings.poll_interval, + ) + executor = PlaywrightExecutor(concurrency=1) + sink = LocalAuthFileSink( + Path(settings.local_auth_dir) / AUTHENTICATED_DIRNAME, + name_secret=settings.source_salt, + ) + inventory = CredentialInventory( + ledger, + Path(settings.local_auth_dir) / AUTHENTICATED_DIRNAME, + Path(settings.local_auth_dir) / CLAIMED_DIRNAME, + ) + recovered = await asyncio.to_thread(inventory.recover) + if recovered: + terminal.emit(("control", {"state": f"recovered {recovered} claims"})) + terminal.emit( + ( + "startup", + { + "destination": f"{AUTHENTICATED_DIRNAME}/", + "source_kind": service_settings.source_kind, + **ledger.inventory_counts(), + }, + ) + ) + pipeline = AuthPipeline( + source=source, + protocol=protocol, + executor=executor, + sink=sink, + ledger=ledger, + timeout=settings.timeout_sec, + min_authorization_interval=( + service_settings.min_authorization_interval_seconds + ), + event_callback=lambda kind, data: terminal.emit((kind, data)), + ) + # Keep base cooldown configurable; the gate grows it after consecutive trips. + pipeline.rate_gate.base_cooldown = float(service_settings.retry_seconds) + pipeline.rate_gate.COOLDOWN_SECONDS = float(service_settings.retry_seconds) + pipeline.rate_gate._current_cooldown = float(service_settings.retry_seconds) + runner = AuthPipelineRunner(pipeline, terminal.emit, inventory=inventory) + await _run_interactive(runner, terminal) + finally: + if pipeline is not None: + pipeline.request_stop() + await client.aclose() + + +def _known_configuration_key(error): + message = str(error) + keys = ( + "XAI_AUTH_SERVICE_LOG_MODE", + "XAI_AUTH_SERVICE_SOURCE", + "XAI_AUTH_SERVICE_SSH_HOST", + "XAI_AUTH_SERVICE_REGISTER_ROOT", + "XAI_AUTH_SERVICE_SYNC_SEC", + "XAI_AUTH_SERVICE_RETRY_SEC", + "XAI_AUTH_SERVICE_MIN_INTERVAL_SEC", + "XAI_ENROLLER_TARGET", + "XAI_ENROLLER_RETRY_ATTEMPTS", + "XAI_ENROLLER_TIMEOUT_SEC", + "XAI_ENROLLER_POLL_SEC", + "XAI_ENROLLER_AUTH_EXECUTOR", + "XAI_ENROLLER_SINK", + "XAI_ENROLLER_CPA_BASE_URL", + "XAI_ENROLLER_CPA_MANAGEMENT_SECRET", + "XAI_ENROLLER_LOCAL_AUTH_DIR", + "XAI_ENROLLER_SOURCE_SALT", + ) + explicit = next((key for key in keys if key in message), None) + if explicit is not None: + return explicit + message_keys = { + "target must be between 1 and 100": "XAI_ENROLLER_TARGET", + "retry attempts must be between 0 and 3": "XAI_ENROLLER_RETRY_ATTEMPTS", + "timeout must be positive": "XAI_ENROLLER_TIMEOUT_SEC", + "poll interval must be positive": "XAI_ENROLLER_POLL_SEC", + "executor must be http or playwright": "XAI_ENROLLER_AUTH_EXECUTOR", + "sink must be cpa or local": "XAI_ENROLLER_SINK", + "CPA base URL must use HTTPS": "XAI_ENROLLER_CPA_BASE_URL", + "CPA base URL and management secret are required": ( + "XAI_ENROLLER_CPA_BASE_URL/XAI_ENROLLER_CPA_MANAGEMENT_SECRET" + ), + "local auth directory is required": "XAI_ENROLLER_LOCAL_AUTH_DIR", + "source salt is required": "XAI_ENROLLER_SOURCE_SALT", + } + return message_keys.get(message) + + +def _print_unexpected_service_failure(error, *, log_mode): + if log_mode == "debug": + detail = f"异常类别 {type(error).__name__}" + else: + detail = "使用 bash auth-service.sh --debug 查看异常类别" + print(f"[✗] 认证服务异常终止 | {detail}", file=sys.stderr) + + +def main(argv=None): + load_dotenv(dotenv_path=Path(__file__).resolve().parents[1] / ".env") + mode = "user" + try: + mode = resolve_auth_log_mode(argv) + asyncio.run(main_async(log_mode=mode)) + return 0 + except AuthServiceConfigurationError as error: + key = _known_configuration_key(error) or "认证服务参数" + print( + f"[✗] 配置错误:{key} | 教程 docs/guides/auth-service.md#配置远端同步", + file=sys.stderr, + ) + return 2 + except ValueError as error: + key = _known_configuration_key(error) + if key is not None: + print( + f"[✗] 配置错误:{key} | 教程 docs/guides/auth-service.md#配置远端同步", + file=sys.stderr, + ) + return 2 + _print_unexpected_service_failure(error, log_mode=mode) + return 1 + except KeyboardInterrupt: + return 130 + except Exception as error: + _print_unexpected_service_failure(error, log_mode=mode) + return 1 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/xai_enroller/sinks.py b/xai_enroller/sinks.py new file mode 100644 index 0000000..539268c --- /dev/null +++ b/xai_enroller/sinks.py @@ -0,0 +1,106 @@ +import asyncio +import hashlib +import hmac +import json +import os +import tempfile +from pathlib import Path +from typing import Protocol +from urllib.parse import urlencode + +import httpx + +from .models import OAuthCredential, SinkReceipt + + +class SinkError(RuntimeError): + pass + + +class CredentialSink(Protocol): + async def store(self, credential: OAuthCredential) -> SinkReceipt: ... + + +def cpa_document(credential: OAuthCredential): + return { + "type": "xai", + "access_token": credential.access_token, + "refresh_token": credential.refresh_token, + "id_token": credential.id_token, + "token_type": credential.token_type, + "expires_in": credential.expires_in, + "expired": credential.expires_at, + "last_refresh": credential.last_refresh, + "sub": credential.subject, + "base_url": "https://api.x.ai/v1", + "token_endpoint": credential.token_endpoint, + "auth_kind": "oauth", + } + + +def credential_filename(credential: OAuthCredential, name_secret: bytes): + subject = credential.subject or credential.refresh_token + digest = hmac.new(name_secret, subject.encode(), hashlib.sha256).hexdigest()[:16] + return f"xai-{digest}.json" + + +class CPAAuthFileSink: + def __init__(self, base_url, management_secret, client: httpx.AsyncClient, name_secret=None): + self.base_url = base_url.rstrip("/") + self.management_secret = management_secret + self.client = client + self.name_secret = name_secret or management_secret.encode() + + async def store(self, credential: OAuthCredential): + filename = credential_filename(credential, self.name_secret) + document = cpa_document(credential) + response = await self.client.post( + f"{self.base_url}/v0/management/auth-files?{urlencode({'name': filename})}", + headers={ + "Authorization": f"Bearer {self.management_secret}", + "Content-Type": "application/json", + }, + json=document, + follow_redirects=False, + ) + if response.status_code // 100 != 2: + raise SinkError("CPA upload rejected") + return SinkReceipt(filename.removesuffix(".json")) + + +class LocalAuthFileSink: + """Atomically persist canonical CPA-compatible documents locally.""" + + def __init__(self, directory, *, name_secret: bytes): + self.directory = Path(directory).expanduser() + self.name_secret = name_secret + + async def store(self, credential: OAuthCredential): + return await asyncio.to_thread(self._store_sync, credential) + + def _store_sync(self, credential: OAuthCredential): + self.directory.mkdir(mode=0o700, parents=True, exist_ok=True) + os.chmod(self.directory, 0o700) + filename = credential_filename(credential, self.name_secret) + destination = self.directory / filename + payload = json.dumps(cpa_document(credential), ensure_ascii=False, indent=2) + "\n" + fd, temporary_name = tempfile.mkstemp( + prefix=f".{filename}.", suffix=".tmp", dir=self.directory, text=True + ) + try: + os.fchmod(fd, 0o600) + with os.fdopen(fd, "w", encoding="utf-8") as stream: + stream.write(payload) + stream.flush() + os.fsync(stream.fileno()) + os.replace(temporary_name, destination) + os.chmod(destination, 0o600) + directory_fd = os.open(self.directory, os.O_RDONLY) + try: + os.fsync(directory_fd) + finally: + os.close(directory_fd) + finally: + if os.path.exists(temporary_name): + os.unlink(temporary_name) + return SinkReceipt(filename.removesuffix(".json")) diff --git a/xai_enroller/sources.py b/xai_enroller/sources.py new file mode 100644 index 0000000..d7598c9 --- /dev/null +++ b/xai_enroller/sources.py @@ -0,0 +1,61 @@ +import base64 +import hashlib +import hmac +import sqlite3 +from pathlib import Path +from typing import Iterator, Protocol + +from .config import require_private_regular_file +from .models import SourceRecord + + +class SourceAdapter(Protocol): + def records(self) -> Iterator[SourceRecord]: ... + + +class FileSourceAdapter: + def __init__(self, path: Path): + self.path = Path(path) + + def records(self) -> Iterator[SourceRecord]: + require_private_regular_file(self.path, require_0600=True) + seen = set() + with self.path.open("r", encoding="utf-8") as handle: + for line_number, line in enumerate(handle, 1): + raw = line.rstrip("\r\n") + if not raw: + continue + source_id, separator, token = raw.partition("\t") + if not separator or not source_id or not token: + raise ValueError(f"invalid source line {line_number}") + if source_id in seen: + continue + seen.add(source_id) + yield SourceRecord(source_id, token) + + +class SQLiteSourceAdapter: + def __init__(self, path: Path, salt: bytes): + self.path = Path(path) + self.salt = bytes(salt) + + def records(self) -> Iterator[SourceRecord]: + require_private_regular_file(self.path) + uri = f"file:{self.path.absolute()}?mode=ro" + connection = sqlite3.connect(uri, uri=True) + try: + rows = connection.execute( + "SELECT token FROM accounts " + "WHERE status = 'active' AND deleted_at IS NULL " + "ORDER BY updated_at DESC" + ) + seen = set() + for (token,) in rows: + fingerprint = hmac.new(self.salt, token.encode(), hashlib.sha256).digest() + source_id = base64.urlsafe_b64encode(fingerprint).rstrip(b"=").decode() + if source_id in seen: + continue + seen.add(source_id) + yield SourceRecord(source_id, token) + finally: + connection.close()