diff --git a/CHANGELOG.md b/CHANGELOG.md index df9d25b..5ba8b05 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,10 @@ ## Unreleased +- 单用户可同时运行的重型任务由 2 个提升到 3 个;任务等待执行容量时,对话会直接说明是当前用户、整机或宿主内存限制,获得槽位后自动继续。 + +- 管理员现在可在管理后台安全录入、测试和更换模型、媒体、检索及语音服务凭据,并查看来源、脱敏尾号和可用状态;数据库凭据可随时删除并回退原有环境配置。DeepSeek 余额低于 30 元、额度耗尽或认证失败时会主动提醒。 + - Mermaid 图表采用更清晰的科研配色与更精致的图框,新生成的流程图还会按数据、处理、判断、结果等角色使用协调的语义色,减少单调的灰白图。 - 管理后台总览按“当前运行”和“运营资源”重新组织,执行容量改为精简摘要并可从右侧详情面板查看队列、容器、用户占用与宿主资源,异常状态和近期指标也更容易识别。 diff --git a/DESIGN.md b/DESIGN.md index 405e8e4..26a6414 100644 --- a/DESIGN.md +++ b/DESIGN.md @@ -257,7 +257,7 @@ scheduled_jobs(§8.5) channel_bindings(§8.7,判别列+JSONB) **边界划分**:Control plane 留宿主(auth/DB/files 校验/SSE/LLM/受控 web 工具/配额审计),Execution plane 进容器(shell/run_python/任意生成代码)。目标不是"所有操作进容器",是"所有不可信执行不能在宿主"——否则凭据反被带进执行面。 -**硬限制**:单容器 4 GiB/2 CPU、pids 1024、`/dev/shm` 512 MiB、`/tmp` 1 GiB、exec timeout、read-only rootfs、no-new-privileges、drop ALL caps、非 root。整机重型执行由宿主共享文件账本 + advisory lock 统一准入:前后台合计 6、后台最多 4、单用户合计 2;蓝绿/多 Web 实例共享同一账本,前台排队不计 timeout 且可取消,宿主 MemAvailable 低于阈值时暂停新放行、不杀存量。轻量 fs 工具不占重型槽。普通容器显式跟踪 active exec,reaper 只回收无 active exec 且超 TTL 的容器。**软配额**:按 user 计 DB(磁盘/LLM cost/wall time/流量),超额 429。**网络**:默认 deny outbound,搜索抓取走宿主受控工具。 +**硬限制**:单容器 4 GiB/2 CPU、pids 1024、`/dev/shm` 512 MiB、`/tmp` 1 GiB、exec timeout、read-only rootfs、no-new-privileges、drop ALL caps、非 root。整机重型执行由宿主共享文件账本 + advisory lock 统一准入:前后台合计 6、后台最多 4、单用户合计 3;蓝绿/多 Web 实例共享同一账本,前台排队不计 timeout 且可取消,并通过 SSE 明示单用户、整机、内存压力或队列顺序等等待原因;宿主 MemAvailable 低于阈值时暂停新放行、不杀存量。轻量 fs 工具不占重型槽。普通容器显式跟踪 active exec,reaper 只回收无 active exec 且超 TTL 的容器。**软配额**:按 user 计 DB(磁盘/LLM cost/wall time/流量),超额 429。**网络**:默认 deny outbound,搜索抓取走宿主受控工具。 **落地清单(Stage C 硬协议,实施按此对账)**: 1. **网络 blocklist 硬编码段**(任一缺失=未完成):`169.254/16`(metadata)、内网三段、CGNAT `100.64/10`;**PG 实际 IP 单独再 block**(belt-and-suspenders)。**容器自身 loopback(`-o lo`)显式放行**——netns 隔离下容器内 127.0.0.1 到不了宿主,DROP 它无安全收益且误伤容器内 IPC(2026-07 实锤:puppeteer↔chromium DevTools 走 127.0.0.1,被 DROP 导致 mermaid 渲染 90 天 0 成功)。 @@ -415,7 +415,7 @@ scheduled_jobs(§8.5) channel_bindings(§8.7,判别列+JSONB) 缺口:markitdown 只抽 PDF 文本层,扫描件(老标准/检测报告/红头指南,建材院高频)转出为空=死路。**选复用 seed-2.0-lite 的方舟文档理解**(chat file 内容块,PDF 整本 base64 内联)新增 `read_document`:零新供应商(敏感文档不出已有豆包面)、零新基础设施、记账复用 vision 通道;实测 ~1300 输入 token/页(约 1 厘/页)、100 页全覆盖、17MB 内联可用。**不选专用解析 API**(MinerU/Textin:版面还原最好,但申报书/专利底稿要上传新第三方 + 免费额度政策不稳);**不选本地 OCR**(PaddleOCR 类:镜像塞推理依赖,需求未量化前过度投资);**不选 file_url/file_id 传址**(前者要给用户文件开免认证公网直链=新安全面、开发机 NAT 后还跑不通;后者要接 TOS 多落一份存储;base64 是零新增面的唯一形态,行业惯例 chat 端点也不收 multipart)。防上下文爆:多页 OCR 强制 `save_md` 落盘只返预览。**升级信号**:>100 页/>30MB 巨件成高频 → 接 TOS 走 file_id;要高保真版面/公式还原 → 再评 MinerU。probe/smoke 留仓(`scripts/probe_ark_doc.py` / `smoke_read_document.py`)。 - **对话锁(前端)**:bg proc 运行期间该 task 的 composer 锁定(发送→停止,Enter 拦截),观感与前台执行完全一致 —— 后台化的收益定位为「进程扛超时/服务重启」,**不改变"一个任务同时只做一件事"的对话心智**;完成的那次轮询解锁 + toast「可继续对话」。锁只在前端,服务端不 409:「停止」入口必须可达,且多设备/渠道绕过前端锁属可接受边缘(等的是同一个进程,发了消息也不冲突)。 -- **防失控**:前后台统一使用整机 6 / 后台 4 / 单用户 2 的共享准入;前台默认超时不放大(它是逼模型做前台/后台选择的杠杆)。 +- **防失控**:前后台统一使用整机 6 / 后台 4 / 单用户 3 的共享准入;前台默认超时不放大(它是逼模型做前台/后台选择的杠杆)。 **边界(防滑坡)**:只覆盖「单个本地长进程」。①**外部异步作业**(seedance 等 submit/poll 形态)不进这里——工具内轮询 + `resume_task_id` 续查已够;②**job 链/依赖/自动重试**不做——那是 workflow 引擎,编排的唯一归属是 agent loop(模型 check 后自己决定下一步),同 §6 拒绝编排的理由;③**完成后自动续跑 run**不做——zcbot 的长任务产物多为终点交付物(与 Claude Code"build 是中间步骤"不同),自动续跑=无人在场烧 token,通知给人、下一步由人/下次对话决定。 @@ -524,6 +524,16 @@ ANSYS 能力固定面向 Windows 上的 Mechanical 2024 R2(revision 242), 当前网关与 zcbot 共用 host,但 HTML 响应同时用 CSP `sandbox` 和 iframe sandbox 强制 opaque origin,且不授予 `allow-same-origin`;这一约束在签名 URL 被顶层打开时同样生效,项目代码不能读取主站 localStorage/JWT。`Referrer-Policy: no-referrer` 防能力 URL 随外链泄露,资源 URL 自带同一签名所以 ES module 不依赖认证 Cookie。将来需要 LocalStorage、Service Worker 或全栈服务器时,升级条件是独立 wildcard preview origin + 专用只挂载当前项目目录的容器 + HTTP/WebSocket/SSE 代理、健康检查和空闲回收,而不是放宽当前同源 sandbox。 +### 8.18 Provider API 凭据动态控制面(implementation,2026-09-02) + +模型、媒体、平台数据与语音 Provider 采用代码内静态可信注册表:Provider id、显示名、凭据字段组、固定测试协议、余额能力和巡检频率均由代码定义,Admin 只能提交凭据值,不能编辑目标 URL、`api_base`、模型、协议或请求模板。这一控制面与可编辑的用户外部系统 definition 分离,避免把平台根凭据送往管理员可变目标。 + +`provider_credentials` 每个 Provider 只保存一套当前凭据。各字段以 `provider + field` 为 AES-GCM AAD 独立加密,未配置 `ZCBOT_CREDENTIAL_MASTER_KEY` 时拒绝写数据库,不存在明文降级。数据库覆盖优先;没有完整覆盖时沿用既有 env,模型能力、价格、`model_id` 和 `api_base` 仍以 `config/models` / `config/media` 为事实源。 + +凭据生命周期采用请求级解析:LLM 在每次构造新的 provider 请求时解析一次,媒体、检索和语音在每次新 HTTP/WebSocket 会话开始时解析;因此更新后下一次请求立即生效,已经发出的 HTTP 或流式响应保持启动时的 Key。候选凭据先在内存中测试,认证成功(含低余额告警)后才以 revision 乐观并发原子替换;失败不改旧密文。 + +Provider 测试将 402/明确额度不足、认证失败、网络不可达和普通 429 分开分类。DeepSeek 读取官方 CNY 余额,`< ¥30` 为低余额,30 分钟巡检;其他 Provider 没有官方余额字段时只展示认证/连通状态。业务调用的明确 402、额度耗尽或认证失败进入同一状态与提醒入口。后台一轮以 PostgreSQL advisory lock 在蓝绿实例间选主,首次异常立即邮件,持续异常最多每日一次,恢复后重置提醒状态。 + --- ## 附录:DeepSeek V4 关键事实(2026-04-24) diff --git a/PROGRESS.md b/PROGRESS.md index 55d9684..452f9eb 100644 --- a/PROGRESS.md +++ b/PROGRESS.md @@ -20,6 +20,10 @@ --- ## 已完成关键能力 +- **09-02 / Unreleased / 单用户重型执行容量与排队提示**:共享执行容量保持整机 6、后台 4,单用户上限由 2 提升到 3;前台重型任务首次等待时通过 SSE 告知单用户、整机、内存压力或队列顺序原因,放行后恢复正常执行提示,排队仍不计入命令 timeout 且可取消。无 schema、migration 或 API 变化。 + +- **09-02 / Unreleased / Admin Provider API 凭据管理与余额提醒**:新增代码内静态可信 Provider 注册表和单行当前凭据控制面,0039 以 AES-GCM 字段级 AAD 密文保存;Admin 可查看 DB/env 来源与尾号、先测后原子替换、手动测试和删除覆盖,LLM/媒体/平台来源/语音均按新请求热解析。DeepSeek 每 30 分钟检查官方 CNY 余额,低于 ¥30 及明确 402/认证失败进入统一状态与邮件提醒,蓝绿实例以 PostgreSQL advisory lock 单轮选主;持续异常每日最多一次、恢复重置。未连接生产数据库、未调用真实第三方 API。 + - **09-02 / Unreleased / 管理后台运行总览分层**:Admin 八张摘要卡按当前运行与运营资源重排,主指标统一为近期或实时口径,累计与低频明细降级;执行容量新增可自动刷新的右侧详情抽屉,集中展示前后台队列、容器生命周期、单用户占用与宿主资源,并移除物理占用与逻辑配额口径不一致的存储进度条。无 API、schema、migration 或运行方式变化。 - **09-02 / 0.70.0 / 管理与更新日志界面收敛**:Admin 顶部“容器状态”直接承载前后台执行、排队、回收、宿主资源与单用户占用,移除重复的独立容量区块,并统一采用“容器依赖”“专业软件节点”用户文案;用户版更新日志只保留可感知变化,公开接口新增向后兼容的 offset 分页元数据,前端每页加载 5 个版本并按需继续加载。 @@ -445,6 +449,7 @@ core/llm_transport.py 438 ← wire 层健壮性:畸形/吐空检测+留 core/tool_registry.py 264 ← 声明式工具注册表((组名,gate,factory);secret/host 工具按实际能力 gate) core/context.py 95 ← LLM 调用前压缩旧 tool / load_skill 消息(带压力门槛),保 tool_call 协议字段 core/external_systems/*.py ← 外部系统目录/用户授权/凭据加密 + 通用 OpenAPI/MCP connector +core/provider_credentials/*.py ← 静态平台 Provider 目录、DB/env resolver、候选测试、状态提醒与蓝绿巡检 core/software_nodes.py ← Windows Node 注册码、身份认证与运行状态 core/sinks.py 101 core/paths.py 50 ← task_dir db form 归一 @@ -465,7 +470,7 @@ core/agent_builder.py 649 ← 装配 lib:build_agent/system prompt(工具 core/executor.py / sandbox/{network,pool,capacity,package_scans}.py / executor_docker.py / procs.py ← Executor ABC + Docker per-user 容器池 + 宿主共享执行槽 + bg proc 文件队列 + 临时依赖扫描 tools/{base,output,fs,shell,run_python,skill_tool,skill_authoring,media_common,seedream,seedance,gpt_image,look_at_image,read_document,image_ref,web_search,web_fetch,documents,materials_project,transcribe_audio,office_to_pdf,external_systems}.py ← media_common=媒体五工具共享原语;external_systems=host-side 外部系统元工具 main.py ~210 ← 入口:web / db / probe / user / sandbox check -db/migrations/versions/ 0001-0038 +db/migrations/versions/ 0001-0039 web/app.py ~210 ← 工厂 + lifespan 编排(07-23 拆分;路由在 routers/,协程在 background 等) web/routers/*.py ← 含 external_systems 用户连接与 software_nodes 节点路由 web/{background,scheduler_runner,wechat_runner}.py ← lifespan 后台协程按域析出 diff --git a/RUN.md b/RUN.md index 7fa791c..88ea26f 100644 --- a/RUN.md +++ b/RUN.md @@ -140,7 +140,8 @@ # ZCBOT_CREDENTIAL_MASTER_KEY=<至少 32 字符随机串> ``` > litellm 在 import 时副作用加载 .env;入口走 `main.py`,`.env` 自动生效。直跑 `python -c "from core.storage import ..."` 不经 litellm 链路时记得自己 `import litellm` 触发,或手动 `export ZCBOT_DB_URL=...`。 -- **平台托管来源配置**:`PAPER_SERVER_*`、`DOCUMENT_SEARCH_*`、`MP_API_KEY` 由宿主部署环境提供,修改后重启 web 生效;不进数据库、用户连接、prompt、tool result、`run_python` 或 Docker sandbox。若需求升级为用户级凭据/per-user grant,接入 `external_systems`;若需要在线动态 definition、revision/reverify 或 OAuth,先设计独立 `platform_sources` 控制面,不直接扩展当前静态 env registry。 +- **平台 Provider 凭据**:`DEEPSEEK_API_KEY`、`ZHIPUAI_API_KEY`、`ARK_API_KEY`、`UNIFYLLM_API_KEY`、`LOCAL_LLM_API_KEY`、`BOCHA_API_KEY`、`PAPER_SERVER_API_KEY`、`DOCUMENT_SEARCH_API_KEY`、`MP_API_KEY` 和两组 `XFYUN_*` 可继续由宿主 env 提供,也可在 0039 migration 后由 Admin「API 凭据」加密覆盖。数据库覆盖在下一次新外部请求生效;删除覆盖立即回退当前 env。URL、模型和协议继续由代码/YAML 管理,Admin 不可编辑。凭据不进入 prompt、tool result、`run_python` 或 Docker sandbox。 +- **Provider 凭据部署**:先配置至少 32 字符的 `ZCBOT_CREDENTIAL_MASTER_KEY`(轮换沿用 `ZCBOT_CREDENTIAL_KEY_ID` / `ZCBOT_CREDENTIAL_PREVIOUS_KEYS`),再执行 `.venv/Scripts/python.exe main.py db upgrade head` 创建 `provider_credentials`。蓝绿部署须先迁移、后启动新代码。未配置 master key 时 env 调用保持可用,但 Admin 保存会明确拒绝;master key 丢失或 AAD 不匹配时已有 DB 覆盖不会静默回退 env,应恢复正确 keyring 或在可解密后删除覆盖。 - **依赖**:`pip install -r requirements.txt`(已在 `.venv` 里;含 `bcrypt`、`segno`、`cryptography`)。 - **微信接入(ClawBot,§8.7)**:① `main.py db upgrade head` 带上 migration `0012`;② `.env` 设 `ZCBOT_WECHAT_BOT_ENABLED=1` + `ZCBOT_WECHAT_SECRET_KEY=<串>`;③ 用户登录后点**左栏 rail「微信」按钮**(`/static/wechat_bind.html` 仍保留作独立/嵌入入口)扫码绑定(需个人微信 8.0.70+ 且灰度到 ClawBot 插件)。绑定后在微信「微信 ClawBot」对话即走 zcbot;**主动推送需用户近 24h 在微信开口过一次**(冷启动/超期推不出,退邮件兜底)。**支持语音消息**(voice_item SILK v3 → pilk 解码 → 讯飞 IAT 转写进对话,回执「🎤 已识别:…」;需 `XFYUN_*` 三件套 + ffmpeg + pilk,pilk 随 requirements 装)。 - **企业微信(渠道 B,纯推送,§8.7)**:① 管理员建自建应用 → 填 `WECOM_CORPID/AGENTID/SECRET`(+ 可见范围含目标用户);② `main.py db upgrade head`。**绑定两条路,任选**: @@ -360,6 +361,10 @@ $env:ZCBOT_EVAL_TOKEN = "" | `GET /v1/skills/{name}` | 返某 skill 完整 SKILL.md 正文(前端「技能」modal 点开查看);同名按 user wins | 必填 | | `DELETE /v1/skills/{name}` | 删当前 user 私有 skill(`.skills//` 整目录);只删 user 源,内置不可删 → 404;`.skills` 文件面板隐藏,这是 UI 上删自己 skill 的唯一入口 | 必填 | | `GET /v1/external-system-providers` | 列管理员已启用的外部系统目录和安全的动态凭据字段声明;不返回完整配置或密钥 | 必填 | +| `GET /v1/admin/provider-credentials` | Admin 查看静态 Provider 目录、来源、脱敏尾号、测试/余额/提醒状态;不返回密文或明文 | admin | +| `PUT /v1/admin/provider-credentials/{provider_id}` | Admin 提交完整凭据组与 `expected_revision`;候选先测试,成功后原子替换 | admin | +| `POST /v1/admin/provider-credentials/{provider_id}/test` | Admin 测试当前 DB/env 凭据;标记 billable 的 Provider 前端会二次确认 | admin | +| `DELETE /v1/admin/provider-credentials/{provider_id}` | Admin 按 revision 删除数据库覆盖并回退 env | admin | | `GET/POST /v1/external-systems` | 列当前用户连接 / 新建并在线验证连接;创建 body `{definition_id,name,credentials}`,旧 Factory `{username,password}` 请求继续兼容 | 必填 | | `PUT /v1/external-systems/{id}/credentials` | 用 `{credentials}` 重新提交并在线验证当前用户连接;凭据不提供读取接口,旧用户名/密码格式继续兼容 | 必填 | | `POST /v1/external-systems/{id}/test` | 用已保存密文凭据测试登录和 Swagger 可读性,并更新连接状态 | 必填 | @@ -391,7 +396,7 @@ $env:ZCBOT_EVAL_TOKEN = "" | `GET /v1/models` | 列 chat LLM 模型清单(扫 `config/models/*.yaml`),前端顶栏切换 / 新建对话框下拉用 | 必填 | | `GET /v1/image_models` | 列图像生成 variant 清单(扫 `config/media/doubao.yaml` image 段),前端"生图"下拉用;yaml 无 image variant → 空列表 → UI 隐藏下拉 | 必填 | -**SSE 事件**(每帧 `event: ` + `data: `):建连时若当前 run 已发布计划,先补 `progress_snapshot{run_id,steps,waiting}` → `run_start{}` → `llm_start{}` → `text{delta}` / `tool_call{name,args,args_preview}` / `tool_result{name,preview,truncated}` → `llm_end{prompt_tokens,completion_tokens}` → `done{}`;cancel 走 `cancelled{}` 后随 `done{}` 收流;异常走 `error{msg}`。`task_progress` 新协议每次携带完整 `steps`,客户端整体替换;消息分页响应也附加 `progress_snapshot`,刷新不依赖当前 30 条窗口。`waiting=true` 表示本轮已调用 `ask_user` 等待确认;正常完成回看时隐藏进度,等待/取消/异常则折叠保留。30s 无 event 服务端发 `: ping` 心跳。nginx 反代记得关 buffering(响应头已带 `X-Accel-Buffering: no` 默认起效)。 +**SSE 事件**(每帧 `event: ` + `data: `):建连时若当前 run 已发布计划,先补 `progress_snapshot{run_id,steps,waiting}` → `run_start{}` → `llm_start{}` → `text{delta}` / `tool_call{name,args,args_preview}` / `execution_queue{state,reason,...}` / `tool_result{name,preview,truncated}` → `llm_end{prompt_tokens,completion_tokens}` → `done{}`;`execution_queue` 仅在前台重型工具确实等待共享容量时发送 `waiting`,放行后发送 `admitted`。cancel 走 `cancelled{}` 后随 `done{}` 收流;异常走 `error{msg}`。`task_progress` 新协议每次携带完整 `steps`,客户端整体替换;消息分页响应也附加 `progress_snapshot`,刷新不依赖当前 30 条窗口。`waiting=true` 表示本轮已调用 `ask_user` 等待确认;正常完成回看时隐藏进度,等待/取消/异常则折叠保留。30s 无 event 服务端发 `: ping` 心跳。nginx 反代记得关 buffering(响应头已带 `X-Accel-Buffering: no` 默认起效)。 **SSE 客户端注意**:浏览器原生 `EventSource` 不支持自定义 header,无法塞 Bearer token。要么 `fetch + ReadableStream` 自解 SSE 帧(dev.html 走的就是这条),要么后端日后加 `?token=...` query(目前不支持,避免 token 进 access log)。 @@ -686,8 +691,8 @@ sudo -u zcbot docker network create zcbot-sandbox-net # ZCBOT_SANDBOX_TMP_SIZE=1g # ZCBOT_MAX_ACTIVE_EXECS=6 # ZCBOT_MAX_BACKGROUND_EXECS=4 -# ZCBOT_MAX_ACTIVE_EXECS_PER_USER=2 -# 三个并发 env 只用于向下收紧,代码硬上限固定为 6 / 4 / 2。 +# ZCBOT_MAX_ACTIVE_EXECS_PER_USER=3 +# 三个并发 env 只用于向下收紧,代码硬上限固定为 6 / 4 / 3。 # ZCBOT_MIN_MEM_AVAILABLE=1g # 低于阈值暂停新放行,不杀正在执行的任务 # PG 实际 IP,逗号分隔。defense-in-depth ── 即便落内网三段(§7.5 #1), # init.sh 再加一遍 DROP 规则。生产部署必填。 diff --git a/core/ark_client.py b/core/ark_client.py index 0e02346..767e07b 100644 --- a/core/ark_client.py +++ b/core/ark_client.py @@ -5,7 +5,6 @@ litellm 不覆盖豆包的图像/视频生成端点,这里自己用 httpx 直调 """ from __future__ import annotations -import os from dataclasses import dataclass from pathlib import Path from typing import Any, Optional @@ -35,6 +34,11 @@ class ArkConfig: api_key: str base_url: str raw: dict # 完整 yaml 内容(便于 caller 按 image/video 子键再取) + api_key_env: str = "" + + def request_api_key(self) -> str: + from core.provider_credentials.runtime import resolve_env_secret + return resolve_env_secret(self.api_key_env) if self.api_key_env else self.api_key @classmethod def load(cls, path: Optional[Path] = None) -> Optional["ArkConfig"]: @@ -50,7 +54,8 @@ class ArkConfig: return None data = yaml.safe_load(p.read_text(encoding="utf-8")) or {} env = data.get("api_key_env") or data.get("ark_api_key_env") or "ARK_API_KEY" - key = os.environ.get(env, "").strip() + from core.provider_credentials.runtime import resolve_env_secret + key = resolve_env_secret(env) if not key: return None base = ( @@ -58,7 +63,7 @@ class ArkConfig: or data.get("ark_base_url") or "https://ark.cn-beijing.volces.com/api/v3" ) - return cls(api_key=key, base_url=str(base).rstrip("/"), raw=data) + return cls(api_key=key, base_url=str(base).rstrip("/"), raw=data, api_key_env=env) class ArkClient: @@ -70,10 +75,11 @@ class ArkClient: def __init__(self, cfg: ArkConfig, timeout_s: float = 60.0) -> None: self.cfg = cfg self.timeout_s = timeout_s + self._api_key = cfg.request_api_key() self._client = httpx.Client( base_url=cfg.base_url, headers={ - "Authorization": f"Bearer {cfg.api_key}", + "Authorization": f"Bearer {self._api_key}", }, timeout=timeout_s, ) @@ -118,17 +124,32 @@ class ArkClient: raise ArkTimeoutError(f"network error calling GET {path}: {e}") from e return self._parse(resp, f"GET {path}") - @staticmethod - def _parse(resp: httpx.Response, label: str) -> dict: + def _parse(self, resp: httpx.Response, label: str) -> dict: if resp.status_code >= 400: + try: + from core.provider_credentials.registry import BY_ENV + from core.provider_credentials.service import record_business_failure + binding = BY_ENV.get(self.cfg.api_key_env) + if binding: + record_business_failure( + binding[0], status_code=resp.status_code, detail=resp.text[:300] + ) + except Exception: + pass # ark 错误 body 一般是 {"error": {"code": ..., "message": ...}};能解就解 try: err = resp.json().get("error") or {} msg = err.get("message") or resp.text[:300] + key = self._api_key + msg = str(msg).replace(key, "***") if key else str(msg) code = err.get("code") or resp.status_code raise ArkError(f"{label} → HTTP {resp.status_code} ({code}): {msg}") except ValueError: - raise ArkError(f"{label} → HTTP {resp.status_code}: {resp.text[:300]}") + detail = resp.text[:300] + key = self._api_key + if key: + detail = detail.replace(key, "***") + raise ArkError(f"{label} → HTTP {resp.status_code}: {detail}") try: return resp.json() except ValueError as e: diff --git a/core/asr_lfasr.py b/core/asr_lfasr.py index d72f747..b95e848 100644 --- a/core/asr_lfasr.py +++ b/core/asr_lfasr.py @@ -18,7 +18,6 @@ import base64 import hashlib import hmac import json -import os import time from typing import Callable, Optional @@ -57,15 +56,15 @@ class LfasrCancelled(LfasrError): def is_configured() -> bool: """LFASR 凭据是否齐 —— agent_builder 据此决定挂不挂 transcribe_audio tool。""" - return all( - (os.getenv(k) or "").strip() - for k in ("XFYUN_APPID", "XFYUN_LFASR_SECRET_KEY") - ) + from core.provider_credentials.runtime import provider_available + return provider_available("xfyun_lfasr") def _load_credentials() -> tuple[str, str]: - appid = (os.getenv("XFYUN_APPID") or "").strip() - secret = (os.getenv("XFYUN_LFASR_SECRET_KEY") or "").strip() + from core.provider_credentials.runtime import resolve_credentials + values = resolve_credentials("xfyun_lfasr").values + appid = values.get("appid", "") + secret = values.get("secret_key", "") if not (appid and secret): raise LfasrNotConfigured( "录音转写未配置:需在 .env 设 XFYUN_APPID / XFYUN_LFASR_SECRET_KEY" @@ -99,6 +98,15 @@ def _post(path: str, params: dict, content: Optional[bytes] = None, headers={"Content-Type": "application/octet-stream"} if content else None, timeout=timeout_s, ) + if resp.status_code >= 400: + try: + from core.provider_credentials.service import record_business_failure + record_business_failure( + "xfyun_lfasr", status_code=resp.status_code, + detail=resp.text[:200], + ) + except Exception: + pass resp.raise_for_status() body = resp.json() except LfasrError: @@ -108,6 +116,14 @@ def _post(path: str, params: dict, content: Optional[bytes] = None, code = str(body.get("code") or "") if code != "000000": desc = body.get("descInfo") or "未知错误" + if code in {"10105", "10106", "10107", "11200"}: + try: + from core.provider_credentials.service import record_business_failure + record_business_failure( + "xfyun_lfasr", status_code=401, detail="讯飞认证失败" + ) + except Exception: + pass raise LfasrError(f"讯飞录音转写失败({code}):{desc}", code=code) return body.get("content") or {} diff --git a/core/asr_xfyun.py b/core/asr_xfyun.py index df2d6b8..d1dd0b1 100644 --- a/core/asr_xfyun.py +++ b/core/asr_xfyun.py @@ -18,7 +18,6 @@ import base64 import hashlib import hmac import json -import os import time from contextlib import suppress from typing import Any @@ -84,9 +83,11 @@ def build_auth_url(api_key: str, api_secret: str) -> str: def _load_credentials() -> tuple[str, str, str]: - appid = (os.getenv("XFYUN_APPID") or "").strip() - api_key = (os.getenv("XFYUN_API_KEY") or "").strip() - api_secret = (os.getenv("XFYUN_API_SECRET") or "").strip() + from core.provider_credentials.runtime import resolve_credentials + values = resolve_credentials("xfyun_iat").values + appid = values.get("appid", "") + api_key = values.get("api_key", "") + api_secret = values.get("api_secret", "") if not (appid and api_key and api_secret): raise XfyunASRNotConfigured( "语音识别未配置:需在 .env 设 XFYUN_APPID / XFYUN_API_KEY / XFYUN_API_SECRET" @@ -94,6 +95,16 @@ def _load_credentials() -> tuple[str, str, str]: return appid, api_key, api_secret +def _report_auth_error(code: object) -> None: + if str(code) not in {"10105", "10106", "10107", "11200"}: + return + try: + from core.provider_credentials.service import record_business_failure + record_business_failure("xfyun_iat", status_code=401, detail="讯飞认证失败") + except Exception: + pass + + class XfyunStream: """流式会话:实时喂 PCM 分片,partial 全文经 on_text 异步回调推出(wpgs 已合并)。 @@ -151,6 +162,7 @@ class XfyunStream: msg = json.loads(await self._ws.recv()) code = msg.get("code") if code: + _report_auth_error(code) hint = _ERR_HINTS.get(code, msg.get("message") or "未知错误") raise XfyunASRError(f"讯飞识别失败({code}):{hint}", code=code) data = msg.get("data") or {} @@ -272,6 +284,7 @@ async def transcribe(pcm: bytes, *, language: str = "zh_cn") -> str: msg = json.loads(raw) code = msg.get("code") if code: + _report_auth_error(code) hint = _ERR_HINTS.get(code, msg.get("message") or "未知错误") raise XfyunASRError(f"讯飞识别失败({code}):{hint}", code=code) data = msg.get("data") or {} diff --git a/core/bocha_client.py b/core/bocha_client.py index a60736c..90223b1 100644 --- a/core/bocha_client.py +++ b/core/bocha_client.py @@ -1,7 +1,6 @@ """博查 (Bocha AI) Web Search API 客户端,共享给 web_search tool。""" from __future__ import annotations -import os from dataclasses import dataclass from pathlib import Path from typing import Optional @@ -22,6 +21,11 @@ class BochaError(RuntimeError): class BochaConfig: api_key: str base_url: str + api_key_env: str = "" + + def request_api_key(self) -> str: + from core.provider_credentials.runtime import resolve_env_secret + return resolve_env_secret(self.api_key_env) if self.api_key_env else self.api_key @classmethod def load(cls, path: Optional[Path] = None) -> Optional["BochaConfig"]: @@ -35,12 +39,14 @@ class BochaConfig: return None data = yaml.safe_load(p.read_text(encoding="utf-8")) or {} env = data.get("bocha_api_key_env") or "BOCHA_API_KEY" - key = os.environ.get(env, "").strip() + from core.provider_credentials.runtime import resolve_env_secret + key = resolve_env_secret(env) if not key: return None return cls( api_key=key, base_url=str(data.get("bocha_base_url") or "https://api.bochaai.com/v1").rstrip("/"), + api_key_env=env, ) @@ -49,10 +55,11 @@ class BochaClient: def __init__(self, cfg: BochaConfig, timeout_s: float = 15.0) -> None: self.cfg = cfg + self._api_key = cfg.request_api_key() self._client = httpx.Client( base_url=cfg.base_url, headers={ - "Authorization": f"Bearer {cfg.api_key}", + "Authorization": f"Bearer {self._api_key}", "Content-Type": "application/json", }, timeout=timeout_s, @@ -73,13 +80,21 @@ class BochaClient: raise BochaError(f"博查网络错误: {e}") from e return self._parse(resp) - @staticmethod - def _parse(resp: httpx.Response) -> dict: + def _parse(self, resp: httpx.Response) -> dict: if resp.status_code >= 400: + try: + from core.provider_credentials.service import record_business_failure + record_business_failure( + "bocha", status_code=resp.status_code, detail=resp.text[:300] + ) + except Exception: + pass try: msg = resp.json().get("message", resp.text[:300]) except ValueError: msg = resp.text[:300] + key = self._api_key + msg = str(msg).replace(key, "***") if key else str(msg) raise BochaError(f"博查 API → HTTP {resp.status_code}: {msg}") try: return resp.json() diff --git a/core/llm.py b/core/llm.py index 83676be..187aaa6 100644 --- a/core/llm.py +++ b/core/llm.py @@ -37,16 +37,39 @@ _REQUEST_TIMEOUT_S = int(os.getenv("ZCBOT_LLM_TIMEOUT_S", "600")) class LLM: - def __init__(self, capabilities: ModelCapabilities) -> None: + def __init__( + self, + capabilities: ModelCapabilities, + *, + credential_resolver: Callable[[str], str] | None = None, + ) -> None: self.caps = capabilities - env_name = capabilities.api_key_env or "DEEPSEEK_API_KEY" - self.api_key = os.environ.get(env_name) + self.env_name = capabilities.api_key_env or "DEEPSEEK_API_KEY" + if credential_resolver is None: + from .provider_credentials.runtime import resolve_env_secret + credential_resolver = resolve_env_secret + self._credential_resolver = credential_resolver + self.api_key = self._credential_resolver(self.env_name) self.api_base = capabilities.api_base or None if not self.api_key: raise RuntimeError( - f"环境变量 {env_name} 未设置,无法调用 {capabilities.model_id}" + f"Provider 凭据 {self.env_name} 未设置,无法调用 {capabilities.model_id}" ) + def _request_api_key(self) -> str: + key = self._credential_resolver(self.env_name) + if not key: + # 非注册表内的测试/私有模型仍保持构造期 env 快照兼容;正式受管 + # Provider 不回退快照,删除 DB 覆盖后应立即回到当前 env 或报未配置。 + from .provider_credentials.registry import BY_ENV + if self.env_name not in BY_ENV: + key = self.api_key + if not key: + raise RuntimeError( + f"Provider 凭据 {self.env_name} 未设置,无法调用 {self.caps.model_id}" + ) + return key + def _build_kwargs( self, messages: List[dict], @@ -64,7 +87,9 @@ class LLM: "model": self.caps.model_id, "messages": provider_messages, "temperature": self.caps.optimal_temperature, - "api_key": self.api_key, + # 请求级解析:已经创建的 LLM 在下一次请求热切新 Key;本次 kwargs + # 构建后保持局部快照,流式迭代期间不会中途换 Key。 + "api_key": self._request_api_key(), "timeout": _REQUEST_TIMEOUT_S, } if self.api_base: @@ -85,6 +110,22 @@ class LLM: kwargs["extra_headers"] = {"anthropic-beta": "prompt-caching-2024-07-31"} return kwargs + def _report_provider_error(self, error: Exception) -> None: + try: + from .provider_credentials.registry import BY_ENV + from .provider_credentials.service import record_business_failure + binding = BY_ENV.get(self.env_name) + if binding is None: + return + raw_status = getattr(error, "status_code", 0) + if hasattr(raw_status, "value"): + raw_status = raw_status.value + record_business_failure( + binding[0], status_code=int(raw_status or 0), detail=str(error) + ) + except Exception: + pass + def chat( self, messages: List[dict], @@ -101,6 +142,7 @@ class LLM: return response except (RateLimitError, APIConnectionError, ServiceUnavailableError, Timeout, APIError) as e: last_err = e + self._report_provider_error(e) if attempt == max_retries - 1: break time.sleep(2 ** attempt) @@ -137,8 +179,7 @@ class LLM: yield from self._chat_stream_direct(kwargs, max_retries=max_retries) - @staticmethod - def _chat_stream_direct(kwargs: dict, *, max_retries: int) -> Iterator[Any]: + def _chat_stream_direct(self, kwargs: dict, *, max_retries: int) -> Iterator[Any]: """无取消调用方的直连路径(probe 等保持原有同步语义)。""" last_err: Optional[Exception] = None @@ -148,6 +189,7 @@ class LLM: break except (RateLimitError, APIConnectionError, ServiceUnavailableError, Timeout, APIError) as e: last_err = e + self._report_provider_error(e) if attempt == max_retries - 1: raise time.sleep(2 ** attempt) @@ -166,8 +208,8 @@ class LLM: except Exception: pass - @staticmethod def _chat_stream_interruptible( + self, kwargs: dict, *, max_retries: int, @@ -206,6 +248,7 @@ class LLM: APIError, ) as e: last_err = e + self._report_provider_error(e) if attempt == max_retries - 1: raise if stop.wait(2 ** attempt): diff --git a/core/provider_credentials/__init__.py b/core/provider_credentials/__init__.py new file mode 100644 index 0000000..1197239 --- /dev/null +++ b/core/provider_credentials/__init__.py @@ -0,0 +1,5 @@ +"""Admin 管理的平台 Provider 凭据控制面。""" + +from .runtime import resolve_credentials, resolve_secret + +__all__ = ["resolve_credentials", "resolve_secret"] diff --git a/core/provider_credentials/monitor.py b/core/provider_credentials/monitor.py new file mode 100644 index 0000000..71759c1 --- /dev/null +++ b/core/provider_credentials/monitor.py @@ -0,0 +1,50 @@ +"""蓝绿单轮选主的 Provider 巡检与开发者邮件。""" +from __future__ import annotations + +import os + +from core.storage import get_engine + +from .registry import get_provider +from .service import due_provider_ids, test_current_credentials +from .testing import TestResult + +_LOCK_SQL = "SELECT pg_try_advisory_lock(31331, 2)" +_UNLOCK_SQL = "SELECT pg_advisory_unlock(31331, 2)" + + +def send_notification(provider_id: str, event: str, result: TestResult) -> None: + provider = get_provider(provider_id) + label = "恢复正常" if event == "recovered" else result.detail + print(f"[provider] {provider_id} {event}: {label}") + email = os.getenv("ZCBOT_DEVELOPER_EMAIL", "").strip() + if not email: + return + try: + from tools.send_email import send_email_smtp, smtp_configured + if smtp_configured(): + send_email_smtp( + email, f"[zcbot] {provider.display_name} {label}", + f"Provider: {provider.display_name}\n状态: {event}\n详情: {result.detail}", + ) + except Exception as exc: # noqa: BLE001 - 告警失败不得影响业务状态 + print(f"[provider] notification failed: {type(exc).__name__}") + + +def run_due_checks() -> int: + engine = get_engine() + with engine.connect() as connection: + claimed = bool(connection.exec_driver_sql(_LOCK_SQL).scalar()) + if not claimed: + return 0 + try: + count = 0 + for provider_id in due_provider_ids(): + try: + test_current_credentials(provider_id, notify=send_notification) + count += 1 + except Exception as exc: # noqa: BLE001 - 单 Provider 隔离 + print(f"[provider] {provider_id} check failed: {type(exc).__name__}") + return count + finally: + connection.exec_driver_sql(_UNLOCK_SQL) diff --git a/core/provider_credentials/registry.py b/core/provider_credentials/registry.py new file mode 100644 index 0000000..ece5a14 --- /dev/null +++ b/core/provider_credentials/registry.py @@ -0,0 +1,83 @@ +"""静态可信 Provider 注册表;网络目标和协议不可由 Admin 修改。""" +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class CredentialField: + name: str + label: str + env: str + secret: bool = True + + +@dataclass(frozen=True) +class ProviderDefinition: + provider_id: str + display_name: str + category: str + fields: tuple[CredentialField, ...] + test_kind: str + test_url: str = "" + test_url_env: str = "" + billable: bool = False + balance_supported: bool = False + check_interval_seconds: int = 86400 + low_balance_threshold: float | None = None + + +def _key(env: str, label: str = "API Key") -> CredentialField: + return CredentialField("api_key", label, env) + + +PROVIDERS: tuple[ProviderDefinition, ...] = ( + ProviderDefinition("deepseek", "DeepSeek", "模型", (_key("DEEPSEEK_API_KEY"),), + "deepseek_balance", "https://api.deepseek.com/user/balance", + balance_supported=True, check_interval_seconds=1800, + low_balance_threshold=30.0), + ProviderDefinition("zhipuai", "智谱开放平台", "模型", (_key("ZHIPUAI_API_KEY"),), + "bearer_get", "https://open.bigmodel.cn/api/paas/v4/models"), + ProviderDefinition("ark", "火山方舟", "模型与媒体", (_key("ARK_API_KEY"),), + "bearer_get", "https://ark.cn-beijing.volces.com/api/v3/models"), + ProviderDefinition("unifyllm", "国际旗舰模型网关", "模型与媒体", (_key("UNIFYLLM_API_KEY"),), + "bearer_get", "https://unifyllm.ai/v1/models"), + ProviderDefinition("local_llm", "内网本地模型", "模型", (_key("LOCAL_LLM_API_KEY"),), + "bearer_get", "http://182.54.21.126:9000/v1/models"), + ProviderDefinition("bocha", "博查搜索", "搜索与平台数据", (_key("BOCHA_API_KEY"),), + "bocha_search", "https://api.bochaai.com/v1/web-search", billable=True), + ProviderDefinition("document_search", "内部材料库", "搜索与平台数据", (_key("DOCUMENT_SEARCH_API_KEY"),), + "bearer_get", "https://ai.ctc-zc.com:8100/api/document_search/list_knowledge_bases", + test_url_env="DOCUMENT_SEARCH_URL"), + ProviderDefinition("paper_server", "论文服务", "搜索与平台数据", (_key("PAPER_SERVER_API_KEY"),), + "query_get", "http://paper.xxhhcty.xyz:8080/api/resm/paper/", + test_url_env="PAPER_SERVER_URL"), + ProviderDefinition("materials_project", "Materials Project", "搜索与平台数据", (_key("MP_API_KEY"),), + "mp_get", "https://api.materialsproject.org/materials/summary/?_limit=1"), + ProviderDefinition( + "xfyun_iat", "讯飞语音听写 IAT", "语音", + (CredentialField("appid", "APPID", "XFYUN_APPID"), + CredentialField("api_key", "API Key", "XFYUN_API_KEY"), + CredentialField("api_secret", "API Secret", "XFYUN_API_SECRET")), + "xfyun_iat", billable=True, + ), + ProviderDefinition( + "xfyun_lfasr", "讯飞录音转写 LFASR", "语音", + (CredentialField("appid", "APPID", "XFYUN_APPID"), + CredentialField("secret_key", "Secret Key", "XFYUN_LFASR_SECRET_KEY")), + "xfyun_lfasr", "https://raasr.xfyun.cn/v2/api/getResult", billable=False, + ), +) + +BY_ID = {provider.provider_id: provider for provider in PROVIDERS} +BY_ENV = { + field.env: (provider.provider_id, field.name) + for provider in PROVIDERS for field in provider.fields +} + + +def get_provider(provider_id: str) -> ProviderDefinition: + try: + return BY_ID[provider_id] + except KeyError as exc: + raise ValueError("unknown provider") from exc diff --git a/core/provider_credentials/runtime.py b/core/provider_credentials/runtime.py new file mode 100644 index 0000000..f15072e --- /dev/null +++ b/core/provider_credentials/runtime.py @@ -0,0 +1,73 @@ +"""每个新外部请求调用一次的数据库优先/env fallback 凭据解析器。""" +from __future__ import annotations + +import os +from dataclasses import dataclass + +from sqlalchemy import select +from sqlalchemy.exc import SQLAlchemyError + +from core.external_systems.crypto import decrypt_secret + +from .registry import BY_ENV, get_provider + + +@dataclass(frozen=True) +class ResolvedCredentials: + provider_id: str + source: str + values: dict[str, str] + + +def _env_values(provider_id: str) -> dict[str, str]: + provider = get_provider(provider_id) + return { + field.name: value + for field in provider.fields + if (value := (os.getenv(field.env) or "").strip()) + } + + +def resolve_credentials(provider_id: str) -> ResolvedCredentials: + provider = get_provider(provider_id) + row = None + try: + from core.storage import session_scope + from core.storage.models import ProviderCredential + with session_scope() as session: + row = session.execute( + select(ProviderCredential).where( + ProviderCredential.provider_id == provider_id + ) + ).scalar_one_or_none() + except (RuntimeError, SQLAlchemyError): + # CLI、migration 前部署或 DB 短暂不可用时保持历史 env 行为。 + pass + if row is not None and all(field.name in row.credentials for field in provider.fields): + # 已存在完整 DB 覆盖时,master key/AAD 错误必须显式失败,不能静默绕回 env。 + values = { + field.name: decrypt_secret( + row.credentials[field.name], aad=f"provider:{provider_id}:{field.name}" + ) + for field in provider.fields + } + return ResolvedCredentials(provider_id, "database", values) + values = _env_values(provider_id) + source = "env" if len(values) == len(provider.fields) else "missing" + return ResolvedCredentials(provider_id, source, values) + + +def resolve_secret(provider_id: str, field: str = "api_key") -> str: + return resolve_credentials(provider_id).values.get(field, "") + + +def resolve_env_secret(env_name: str) -> str: + binding = BY_ENV.get(env_name) + if binding is None: + return (os.getenv(env_name) or "").strip() + return resolve_secret(*binding) + + +def provider_available(provider_id: str) -> bool: + provider = get_provider(provider_id) + return len(resolve_credentials(provider_id).values) == len(provider.fields) diff --git a/core/provider_credentials/service.py b/core/provider_credentials/service.py new file mode 100644 index 0000000..9d21a93 --- /dev/null +++ b/core/provider_credentials/service.py @@ -0,0 +1,263 @@ +"""Provider 凭据 CRUD、乐观替换、状态持久化与定时检查。""" +from __future__ import annotations + +from collections.abc import Callable +from datetime import datetime, timedelta, timezone +from uuid import UUID + +from sqlalchemy import delete, select, update +from sqlalchemy.exc import IntegrityError + +from core.external_systems.crypto import configured as crypto_configured +from core.external_systems.crypto import encrypt_secret +from core.storage import session_scope +from core.storage.models import ProviderCredential + +from .registry import PROVIDERS, get_provider +from .runtime import _env_values, resolve_credentials +from .testing import TestResult, classify_response, test_provider + +_BAD = {"low_balance", "exhausted", "auth_error", "unreachable"} +_NOTIFY_COOLDOWN = timedelta(hours=24) + + +class ProviderCredentialError(RuntimeError): + pass + + +class RevisionConflict(ProviderCredentialError): + pass + + +def _hint(value: str) -> str: + value = str(value or "") + return f"***{value[-4:]}" if len(value) >= 4 else "***" + + +def _row_payload(provider, row: ProviderCredential | None) -> dict: + database = bool(row and all(field.name in row.credentials for field in provider.fields)) + env_values = _env_values(provider.provider_id) + env_source = "env" if len(env_values) == len(provider.fields) else "missing" + source = "database" if database else env_source + hints = row.credential_hint if database else { + field.name: _hint(env_values.get(field.name, "")) + for field in provider.fields if env_values.get(field.name) + } + return { + "provider_id": provider.provider_id, + "display_name": provider.display_name, + "category": provider.category, + "fields": [ + {"name": field.name, "label": field.label, "required": True, + "hint": hints.get(field.name, "")} + for field in provider.fields + ], + "configured": source != "missing", + "source": source, + "revision": row.revision if row else 0, + "test_status": row.test_status if row else "untested", + "test_detail": row.test_detail if row else None, + "last_tested_at": row.last_tested_at.isoformat() if row and row.last_tested_at else None, + "balance": ( + {"amount": str(row.balance_amount), "currency": row.balance_currency} + if row and row.balance_amount is not None else None + ), + "balance_supported": provider.balance_supported, + "billable_test": provider.billable, + "alerting": bool(row and row.test_status in _BAD), + } + + +def list_providers() -> list[dict]: + with session_scope() as session: + rows = { + row.provider_id: row + for row in session.execute(select(ProviderCredential)).scalars() + } + return [_row_payload(provider, rows.get(provider.provider_id)) for provider in PROVIDERS] + + +def _validate_values(provider_id: str, values: dict[str, str]) -> dict[str, str]: + provider = get_provider(provider_id) + expected = {field.name for field in provider.fields} + cleaned = {str(k): str(v).strip() for k, v in values.items()} + if set(cleaned) != expected or any(not value for value in cleaned.values()): + raise ProviderCredentialError("凭据字段不完整或包含未知字段") + return cleaned + + +def _notification_transition( + old_status: str, old_notified_at: datetime | None, result: TestResult, now: datetime +) -> str | None: + if result.status == "normal" and old_status in _BAD: + return "recovered" + if result.status not in _BAD: + return None + if old_status not in _BAD or old_status != result.status: + return result.status + if old_notified_at is None or now - old_notified_at >= _NOTIFY_COOLDOWN: + return result.status + return None + + +def _apply_result( + provider_id: str, + result: TestResult, + *, + notify: Callable[[str, str, TestResult], None] | None = None, +) -> None: + now = datetime.now(timezone.utc) + event: str | None = None + with session_scope() as session: + row = session.execute( + select(ProviderCredential).where(ProviderCredential.provider_id == provider_id) + ).scalar_one_or_none() + if row is None: + row = ProviderCredential( + provider_id=provider_id, credentials={}, credential_hint={}, revision=1 + ) + session.add(row) + session.flush() + old_status = row.test_status + event = _notification_transition(old_status, row.last_notified_at, result, now) + row.test_status = result.status + row.test_detail = result.detail[:500] + row.balance_amount = result.balance_amount + row.balance_currency = result.balance_currency + row.last_tested_at = now + if event == "recovered": + row.last_notified_at = None + row.last_notified_status = None + elif event: + row.last_notified_at = now + row.last_notified_status = result.status + if event and notify: + notify(provider_id, event, result) + + +def replace_credentials( + provider_id: str, + values: dict[str, str], + *, + expected_revision: int, + updated_by: UUID, + request=None, + notify: Callable[[str, str, TestResult], None] | None = None, +) -> dict: + if not crypto_configured(): + raise ProviderCredentialError( + "未配置 ZCBOT_CREDENTIAL_MASTER_KEY,禁止保存数据库凭据" + ) + provider = get_provider(provider_id) + cleaned = _validate_values(provider_id, values) + kwargs = {"request": request} if request is not None else {} + result = test_provider(provider_id, cleaned, **kwargs) + if not result.accepted: + raise ProviderCredentialError(f"候选凭据测试失败:{result.detail}") + encrypted = { + field.name: encrypt_secret( + cleaned[field.name], aad=f"provider:{provider_id}:{field.name}" + ) + for field in provider.fields + } + hints = {field.name: _hint(cleaned[field.name]) for field in provider.fields} + now = datetime.now(timezone.utc) + try: + with session_scope() as session: + if expected_revision == 0: + session.add(ProviderCredential( + provider_id=provider_id, credentials=encrypted, + credential_hint=hints, revision=1, test_status=result.status, + test_detail=result.detail, balance_amount=result.balance_amount, + balance_currency=result.balance_currency, last_tested_at=now, + updated_by=updated_by, + )) + revision = 1 + else: + changed = session.execute( + update(ProviderCredential) + .where(ProviderCredential.provider_id == provider_id) + .where(ProviderCredential.revision == expected_revision) + .values( + credentials=encrypted, credential_hint=hints, + revision=expected_revision + 1, test_status=result.status, + test_detail=result.detail, balance_amount=result.balance_amount, + balance_currency=result.balance_currency, last_tested_at=now, + updated_by=updated_by, updated_at=now, + ) + ) + if int(changed.rowcount or 0) != 1: + raise RevisionConflict("凭据已被其他管理员更新,请刷新后重试") + revision = expected_revision + 1 + except IntegrityError as exc: + raise RevisionConflict("凭据已被其他管理员更新,请刷新后重试") from exc + if result.status == "low_balance" and notify: + notify(provider_id, result.status, result) + with session_scope() as session: + session.execute( + update(ProviderCredential) + .where(ProviderCredential.provider_id == provider_id) + .where(ProviderCredential.revision == revision) + .values(last_notified_at=now, last_notified_status=result.status) + ) + return {"provider_id": provider_id, "revision": revision, + "test_status": result.status, "test_detail": result.detail} + + +def test_current_credentials( + provider_id: str, *, request=None, + notify: Callable[[str, str, TestResult], None] | None = None, +) -> TestResult: + resolved = resolve_credentials(provider_id) + provider = get_provider(provider_id) + if len(resolved.values) != len(provider.fields): + raise ProviderCredentialError("Provider 凭据未完整配置") + kwargs = {"request": request} if request is not None else {} + result = test_provider(provider_id, resolved.values, **kwargs) + _apply_result(provider_id, result, notify=notify) + return result + + +def delete_override(provider_id: str, *, expected_revision: int) -> None: + get_provider(provider_id) + with session_scope() as session: + changed = session.execute( + delete(ProviderCredential) + .where(ProviderCredential.provider_id == provider_id) + .where(ProviderCredential.revision == expected_revision) + ) + if int(changed.rowcount or 0) != 1: + raise RevisionConflict("凭据已被更新或不存在,请刷新后重试") + + +def record_business_failure( + provider_id: str, *, status_code: int = 0, detail: str = "", + notify: Callable[[str, str, TestResult], None] | None = None, +) -> str | None: + result = classify_response(status_code, detail) + if result.status not in {"auth_error", "exhausted"}: + return None + if notify is None: + from .monitor import send_notification + notify = send_notification + _apply_result(provider_id, result, notify=notify) + return result.status + + +def due_provider_ids(now: datetime | None = None) -> list[str]: + now = now or datetime.now(timezone.utc) + with session_scope() as session: + rows = { + row.provider_id: row for row in session.execute(select(ProviderCredential)).scalars() + } + due = [] + for provider in PROVIDERS: + resolved = resolve_credentials(provider.provider_id) + if len(resolved.values) != len(provider.fields): + continue + row = rows.get(provider.provider_id) + if row is None or row.last_tested_at is None or ( + now - row.last_tested_at >= timedelta(seconds=provider.check_interval_seconds) + ): + due.append(provider.provider_id) + return due diff --git a/core/provider_credentials/testing.py b/core/provider_credentials/testing.py new file mode 100644 index 0000000..a731648 --- /dev/null +++ b/core/provider_credentials/testing.py @@ -0,0 +1,172 @@ +"""Provider 候选凭据测试与统一、保守的错误分类。""" +from __future__ import annotations + +import base64 +import hashlib +import hmac +import json +import os +import re +import time +from collections.abc import Callable +from dataclasses import dataclass +from decimal import Decimal, InvalidOperation +from typing import Any + +import httpx + +from .registry import ProviderDefinition, get_provider + +_EXHAUSTED_RE = re.compile( + r"insufficient[ _-]?(?:balance|quota|credit)|balance[ _-]?not[ _-]?enough|余额不足|额度不足|quota exhausted", + re.IGNORECASE, +) +_AUTH_RE = re.compile( + r"invalid.*(?:key|token)|authentication|unauthori[sz]ed|鉴权|认证失败", + re.IGNORECASE, +) + + +@dataclass(frozen=True) +class TestResult: + status: str + detail: str + balance_amount: Decimal | None = None + balance_currency: str | None = None + + @property + def accepted(self) -> bool: + return self.status in {"normal", "low_balance"} + + +def classify_response(status_code: int, text: str = "") -> TestResult: + safe = re.sub(r"\s+", " ", text or "").strip()[:300] + if status_code in {401, 403} or _AUTH_RE.search(safe): + return TestResult("auth_error", f"认证失败(HTTP {status_code})") + if status_code == 402 or _EXHAUSTED_RE.search(safe): + return TestResult("exhausted", f"余额或额度已耗尽(HTTP {status_code})") + if status_code == 429: + return TestResult("unreachable", "服务限流(HTTP 429),未判定为余额耗尽") + if status_code >= 400: + return TestResult("unreachable", f"服务返回 HTTP {status_code}") + return TestResult("normal", "认证与连通性正常") + + +def _deepseek(response: httpx.Response, provider: ProviderDefinition) -> TestResult: + base = classify_response(response.status_code, response.text) + if base.status != "normal": + return base + try: + body = response.json() + infos = body.get("balance_infos") or [] + cny = next(x for x in infos if str(x.get("currency", "")).upper() == "CNY") + amount = Decimal(str(cny.get("total_balance"))) + except (ValueError, KeyError, StopIteration, InvalidOperation, TypeError): + return TestResult("unreachable", "余额响应格式异常") + if body.get("is_available") is False or amount <= 0: + return TestResult("exhausted", "余额已耗尽", amount, "CNY") + threshold = Decimal(str(provider.low_balance_threshold or 0)) + status = "low_balance" if amount < threshold else "normal" + detail = f"CNY 可用余额 ¥{amount:.2f}" + return TestResult(status, detail, amount, "CNY") + + +def _xfyun_sign(credentials: dict[str, str]) -> dict[str, str]: + ts = str(int(time.time())) + md5hex = hashlib.md5( + (credentials["appid"] + ts).encode(), usedforsecurity=False + ).hexdigest() + signa = base64.b64encode( + hmac.new(credentials["secret_key"].encode(), md5hex.encode(), hashlib.sha1).digest() + ).decode() + return {"appId": credentials["appid"], "ts": ts, "signa": signa, + "orderId": "zcbot-credential-check", "resultType": "transfer"} + + +def test_provider( + provider_id: str, + credentials: dict[str, str], + *, + request: Callable[..., httpx.Response] = httpx.request, +) -> TestResult: + provider = get_provider(provider_id) + expected = {field.name for field in provider.fields} + if set(credentials) != expected or any(not str(v).strip() for v in credentials.values()): + return TestResult("auth_error", "凭据字段不完整") + headers: dict[str, str] = {} + params: dict[str, str] = {} + json_body: dict[str, Any] | None = None + method = "GET" + test_url = provider.test_url + if provider.test_url_env: + base = (os.getenv(provider.test_url_env) or "").strip().rstrip("/") + if base and provider_id == "document_search": + test_url = f"{base}/document_search/list_knowledge_bases" + elif base and provider_id == "paper_server": + test_url = f"{base}/api/resm/paper/" + if provider.test_kind in {"deepseek_balance", "bearer_get", "bocha_search"}: + headers["Authorization"] = f"Bearer {credentials['api_key']}" + if provider.test_kind == "query_get": + params["api_key"] = credentials["api_key"] + params["page_size"] = "1" + elif provider.test_kind == "mp_get": + headers["X-API-KEY"] = credentials["api_key"] + elif provider.test_kind == "bocha_search": + method = "POST" + json_body = {"query": "水泥", "count": 1, "freshness": "noLimit"} + elif provider.test_kind == "xfyun_lfasr": + method = "POST" + params.update(_xfyun_sign(credentials)) + elif provider.test_kind == "xfyun_iat": + # IAT 鉴权只存在于 WebSocket upgrade;复用官方签名函数并允许测试替换 + # requester。200/101 均视为握手成功,真实 handler 不发送音频、不计费。 + from core.asr_xfyun import build_auth_url + url = build_auth_url(credentials["api_key"], credentials["api_secret"]) + if request is httpx.request: + try: + from websockets.sync.client import connect + with connect(url, open_timeout=8) as websocket: + websocket.send(json.dumps({ + "common": {"app_id": credentials["appid"]}, + "business": {"language": "zh_cn", "domain": "iat", "accent": "mandarin"}, + "data": {"status": 2, "format": "audio/L16;rate=16000", + "encoding": "raw", "audio": ""}, + })) + body = json.loads(websocket.recv(timeout=8)) + if int(body.get("code") or 0) == 0: + return TestResult("normal", "WebSocket 凭据组鉴权正常") + return TestResult("auth_error", "WebSocket 凭据组认证失败") + except Exception as exc: # noqa: BLE001 - WebSocket 库异常族随版本变化 + text = str(exc) + status = int(getattr(exc, "status_code", 0) or 0) + return classify_response(status, text) if status else TestResult( + "unreachable", "WebSocket 连接失败" + ) + response = request("GET", url, headers={"X-Appid": credentials["appid"]}, timeout=8) + return classify_response(response.status_code, response.text) + try: + response = request( + method, test_url, headers=headers, params=params, + json=json_body, timeout=8, + ) + except httpx.HTTPError: + return TestResult("unreachable", "网络连接失败") + if provider.test_kind == "deepseek_balance": + return _deepseek(response, provider) + if provider.test_kind == "xfyun_lfasr" and response.status_code < 400: + try: + body = response.json() + code = str(body.get("code") or "") + description = str(body.get("descInfo") or "") + if code in {"10105", "10106", "10107"}: + return TestResult("auth_error", "认证失败") + # 虚拟订单的“订单不存在”说明签名已通过。 + if code and re.search( + r"order|订单.*(?:不存在|无效)", description, re.IGNORECASE + ): + return TestResult("normal", "认证正常(未创建计费订单)") + if code: + return TestResult("unreachable", "服务未确认凭据有效") + except (ValueError, AttributeError): + pass + return classify_response(response.status_code, response.text) diff --git a/core/storage/models.py b/core/storage/models.py index 563c7e2..0f1642a 100644 --- a/core/storage/models.py +++ b/core/storage/models.py @@ -258,6 +258,48 @@ class WebPreview(Base): ) +class ProviderCredential(Base): + """平台级可信 Provider 当前凭据与健康状态;credentials 只保存字段密文。""" + + __tablename__ = "provider_credentials" + + provider_id: Mapped[str] = mapped_column(Text, primary_key=True) + credentials: Mapped[dict[str, Any]] = mapped_column(JSONB, nullable=False) + credential_hint: Mapped[dict[str, Any]] = mapped_column( + JSONB, nullable=False, default=dict + ) + revision: Mapped[int] = mapped_column( + Integer, nullable=False, default=1, server_default="1" + ) + test_status: Mapped[str] = mapped_column( + Text, nullable=False, default="untested", server_default="untested" + ) + test_detail: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + balance_amount: Mapped[Optional[Decimal]] = mapped_column( + Numeric(18, 6), nullable=True + ) + balance_currency: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + last_tested_at: Mapped[Optional[datetime]] = mapped_column( + DateTime(timezone=True), nullable=True + ) + last_notified_at: Mapped[Optional[datetime]] = mapped_column( + DateTime(timezone=True), nullable=True + ) + last_notified_status: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + updated_by: Mapped[Optional[UUID]] = mapped_column( + PG_UUID(as_uuid=True), + ForeignKey("users.user_id", ondelete="SET NULL"), + nullable=True, + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), server_default=func.now(), nullable=False + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), server_default=func.now(), + onupdate=func.now(), nullable=False, + ) + + class Artifact(Base): """Stable identity and lifecycle metadata for a published workspace file.""" diff --git a/db/migrations/versions/20260902_1000_0039_provider_credentials.py b/db/migrations/versions/20260902_1000_0039_provider_credentials.py new file mode 100644 index 0000000..03e9eda --- /dev/null +++ b/db/migrations/versions/20260902_1000_0039_provider_credentials.py @@ -0,0 +1,46 @@ +"""Add encrypted provider credentials control plane. + +Revision ID: 0039 +Revises: 0038 +Create Date: 2026-09-02 +""" +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op +from sqlalchemy.dialects import postgresql + +revision: str = "0039" +down_revision: str | None = "0038" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "provider_credentials", + sa.Column("provider_id", sa.Text(), primary_key=True), + sa.Column("credentials", postgresql.JSONB(), nullable=False), + sa.Column("credential_hint", postgresql.JSONB(), nullable=False), + sa.Column("revision", sa.Integer(), server_default="1", nullable=False), + sa.Column("test_status", sa.Text(), server_default="untested", nullable=False), + sa.Column("test_detail", sa.Text(), nullable=True), + sa.Column("balance_amount", sa.Numeric(18, 6), nullable=True), + sa.Column("balance_currency", sa.Text(), nullable=True), + sa.Column("last_tested_at", sa.DateTime(timezone=True), nullable=True), + sa.Column("last_notified_at", sa.DateTime(timezone=True), nullable=True), + sa.Column("last_notified_status", sa.Text(), nullable=True), + sa.Column( + "updated_by", + postgresql.UUID(as_uuid=True), + sa.ForeignKey("users.user_id", ondelete="SET NULL"), + nullable=True, + ), + sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False), + sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False), + sa.CheckConstraint("revision > 0", name="ck_provider_credentials_revision_positive"), + ) + + +def downgrade() -> None: + op.drop_table("provider_credentials") diff --git a/platform_sources/materials_library_client.py b/platform_sources/materials_library_client.py index fe2c12e..b069e14 100644 --- a/platform_sources/materials_library_client.py +++ b/platform_sources/materials_library_client.py @@ -39,7 +39,8 @@ _LIST_FIELDS = ( def _api_key() -> str: - key = os.environ.get("DOCUMENT_SEARCH_API_KEY", "").strip() + from core.provider_credentials.runtime import resolve_secret + key = resolve_secret("document_search") if not key: raise RuntimeError( "DOCUMENT_SEARCH_API_KEY env 未设置 —— 配置后再使用 literature skill 的内部材料库来源" @@ -54,6 +55,19 @@ def _auth_headers(extra: Optional[dict] = None) -> dict: return h +def _check_response(response: httpx.Response) -> None: + if response.status_code >= 400: + try: + from core.provider_credentials.service import record_business_failure + record_business_failure( + "document_search", status_code=response.status_code, + detail=response.text[:300], + ) + except Exception: + pass + response.raise_for_status() + + def _safe_name(name: str) -> str: # 防目录穿越;保留扩展名 return name.replace("/", "_").replace("\\", "_").replace("..", "_") @@ -66,7 +80,7 @@ def list_kb() -> list[dict]: / create_time / file_count。只返回 ID 映射里有效的(分类 1-7)。 """ r = httpx.get(f"{_API}/list_knowledge_bases", headers=_auth_headers(), timeout=_TIMEOUT) - r.raise_for_status() + _check_response(r) payload = r.json() data = payload.get("data") or {} return list(data.get("knowledge_bases") or []) @@ -104,7 +118,7 @@ def search( json=body, timeout=_TIMEOUT, ) - r.raise_for_status() + _check_response(r) payload = r.json() data = payload.get("data") or {} docs = data.get("documents") or [] @@ -144,7 +158,7 @@ def download( params=params, timeout=_DOWNLOAD_TIMEOUT, ) as resp: - resp.raise_for_status() + _check_response(resp) with open(dest, "wb") as f: for chunk in resp.iter_bytes(chunk_size=64 * 1024): f.write(chunk) diff --git a/platform_sources/materials_project.py b/platform_sources/materials_project.py index a168dbb..7c0d1d6 100644 --- a/platform_sources/materials_project.py +++ b/platform_sources/materials_project.py @@ -30,7 +30,8 @@ _DEFAULT_SUMMARY_FIELDS = [ def _mp_key() -> str: - key = os.environ.get("MP_API_KEY", "").strip() + from core.provider_credentials.runtime import resolve_secret + key = resolve_secret("materials_project") if not key: raise RuntimeError("MP_API_KEY env 未设置,无法查询 Materials Project") return key @@ -61,6 +62,17 @@ def _mpr(): return MPRester(_mp_key()) +def _report_provider_error(error: Exception) -> None: + try: + from core.provider_credentials.service import record_business_failure + status = int(getattr(error, "status_code", 0) or 0) + record_business_failure( + "materials_project", status_code=status, detail=str(error) + ) + except Exception: + pass + + class MaterialsProjectSearchTool(Tool): name = "materials_project_search" description = ( @@ -128,6 +140,7 @@ class MaterialsProjectSearchTool(Tool): num_chunks=1, chunk_size=limit, **kwargs ) except Exception as e: + _report_provider_error(e) detail = safe_error_text(e, ("MP_API_KEY",)) return f"[Error] materials_project_search failed: {type(e).__name__}: {detail}" plain = [_to_plain(d) for d in list(docs)[:limit]] @@ -152,6 +165,7 @@ class MaterialsProjectSearchTool(Tool): try: session = _mpr() except Exception as e: + _report_provider_error(e) detail = safe_error_text(e, ("MP_API_KEY",)) return f"[Error] materials_project_search batch failed: {type(e).__name__}: {detail}" agg: list[dict[str, Any]] = [] @@ -169,6 +183,7 @@ class MaterialsProjectSearchTool(Tool): n_ok += 1 agg.append({"formula": f, "n": len(results), "results": results}) except Exception as e: + _report_provider_error(e) detail = safe_error_text(e, ("MP_API_KEY",)) agg.append({"formula": f, "error": f"{type(e).__name__}: {detail}"}) return self._render_batch(agg, fields, n_ok) @@ -251,6 +266,7 @@ class MaterialsProjectGetStructureTool(Tool): dest.parent.mkdir(parents=True, exist_ok=True) struct.to(filename=str(dest)) except Exception as e: + _report_provider_error(e) detail = safe_error_text(e, ("MP_API_KEY",)) return f"[Error] materials_project_get_structure failed: {type(e).__name__}: {detail}" return f"saved: {self._display(dest)}" diff --git a/platform_sources/paper_server.py b/platform_sources/paper_server.py index 6c84c9c..a819438 100644 --- a/platform_sources/paper_server.py +++ b/platform_sources/paper_server.py @@ -47,7 +47,8 @@ _LIST_FIELDS = ( def _config() -> tuple[str, str, str]: base_url = os.environ.get("PAPER_SERVER_URL", _DEFAULT_BASE_URL).strip().rstrip("/") - api_key = os.environ.get("PAPER_SERVER_API_KEY", "").strip() + from core.provider_credentials.runtime import resolve_secret + api_key = resolve_secret("paper_server") if not api_key: raise RuntimeError("PAPER_SERVER_API_KEY env 未设置,无法查询 paper_server") return base_url, f"{base_url}/api/resm/paper", api_key @@ -70,6 +71,14 @@ def _raise_response_error(response: httpx.Response) -> None: except Exception: pass if response.status_code in (401, 403) or err_code in _AUTH_ERR_CODES: + try: + from core.provider_credentials.service import record_business_failure + record_business_failure( + "paper_server", status_code=response.status_code or 401, + detail=err_code, + ) + except Exception: + pass raise RuntimeError( f"paper_server auth failed (HTTP {response.status_code}, " f"{err_code or 'no err_code'}):请管理员检查平台 PAPER_SERVER_API_KEY" diff --git a/platform_sources/registry.py b/platform_sources/registry.py index c31a87b..98ef63d 100644 --- a/platform_sources/registry.py +++ b/platform_sources/registry.py @@ -54,7 +54,8 @@ class PlatformSourceProvider(Protocol): def _env_set(name: str) -> bool: - return bool(os.environ.get(name, "").strip()) + from core.provider_credentials.runtime import resolve_env_secret + return bool(resolve_env_secret(name)) def _require_valid_http_url(name: str, default: str) -> None: diff --git a/tests/test_provider_credentials.py b/tests/test_provider_credentials.py new file mode 100644 index 0000000..da8fe68 --- /dev/null +++ b/tests/test_provider_credentials.py @@ -0,0 +1,274 @@ +import json +import os +import unittest +from contextlib import contextmanager +from datetime import datetime, timedelta, timezone +from importlib import util +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock, Mock, patch +from uuid import uuid4 + +import httpx + +from core.capabilities import ModelCapabilities +from core.llm import LLM +from core.provider_credentials.registry import BY_ID +from core.provider_credentials.runtime import resolve_credentials +from core.provider_credentials.service import ( + ProviderCredentialError, + RevisionConflict, + _hint, + _notification_transition, + _row_payload, + replace_credentials, +) +from core.provider_credentials.testing import ( + TestResult, + classify_response, + test_provider, +) + + +def response(status, body): + content = json.dumps(body).encode() if isinstance(body, dict) else str(body).encode() + return httpx.Response(status, content=content) + + +class ProviderRegistryTests(unittest.TestCase): + def test_credential_groups_are_explicit(self): + self.assertEqual( + [f.name for f in BY_ID["xfyun_iat"].fields], + ["appid", "api_key", "api_secret"], + ) + self.assertEqual( + [f.name for f in BY_ID["xfyun_lfasr"].fields], + ["appid", "secret_key"], + ) + self.assertNotIn("ZCBOT_DB_URL", [f.env for p in BY_ID.values() for f in p.fields]) + + def test_hint_only_reveals_last_four(self): + self.assertEqual(_hint("sk-abcdefgh"), "***efgh") + self.assertNotIn("abcdef", _hint("sk-abcdefgh")) + + +class ProviderClassificationTests(unittest.TestCase): + def test_http_402_and_429_are_not_conflated(self): + self.assertEqual(classify_response(402, "").status, "exhausted") + self.assertEqual(classify_response(429, "rate limit").status, "unreachable") + + def test_deepseek_30_yuan_boundary(self): + def request_at(amount): + return lambda *a, **k: response(200, { + "is_available": True, + "balance_infos": [{"currency": "CNY", "total_balance": str(amount)}], + }) + creds = {"api_key": "candidate-secret"} + self.assertEqual( + test_provider("deepseek", creds, request=request_at("29.99")).status, + "low_balance", + ) + self.assertEqual( + test_provider("deepseek", creds, request=request_at("30.00")).status, + "normal", + ) + + def test_candidate_secret_is_not_returned_in_detail(self): + secret = "candidate-super-secret" + result = test_provider( + "deepseek", {"api_key": secret}, + request=lambda *a, **k: response(401, {"message": secret}), + ) + self.assertEqual(result.status, "auth_error") + self.assertNotIn(secret, result.detail) + + +class ProviderResolutionTests(unittest.TestCase): + def test_database_precedes_environment(self): + row = SimpleNamespace(credentials={"api_key": "cipher"}) + + class Result: + def scalar_one_or_none(self): + return row + + class Session: + def execute(self, _statement): + return Result() + + @contextmanager + def scope(): + yield Session() + + with ( + patch.dict(os.environ, {"DEEPSEEK_API_KEY": "env-key"}), + patch("core.storage.session_scope", scope), + patch("core.provider_credentials.runtime.decrypt_secret", return_value="db-key") as decrypt, + ): + resolved = resolve_credentials("deepseek") + self.assertEqual((resolved.source, resolved.values["api_key"]), ("database", "db-key")) + decrypt.assert_called_once_with("cipher", aad="provider:deepseek:api_key") + + def test_environment_fallback_without_database(self): + with ( + patch.dict(os.environ, {"DEEPSEEK_API_KEY": "env-key"}), + patch("core.storage.session_scope", side_effect=RuntimeError("no db")), + ): + resolved = resolve_credentials("deepseek") + self.assertEqual((resolved.source, resolved.values), ("env", {"api_key": "env-key"})) + + def test_each_xfyun_group_resolves_its_declared_fields(self): + env = { + "XFYUN_APPID": "app", "XFYUN_API_KEY": "iat-key", + "XFYUN_API_SECRET": "iat-secret", "XFYUN_LFASR_SECRET_KEY": "lf-key", + } + with patch.dict(os.environ, env), patch("core.storage.session_scope", side_effect=RuntimeError("no db")): + self.assertEqual(set(resolve_credentials("xfyun_iat").values), {"appid", "api_key", "api_secret"}) + self.assertEqual(set(resolve_credentials("xfyun_lfasr").values), {"appid", "secret_key"}) + + +class LLMCredentialLifecycleTests(unittest.TestCase): + def _caps(self): + return ModelCapabilities( + family="deepseek_v4", model_id="deepseek/model", + api_key_env="DEEPSEEK_API_KEY", thinking_transport="none", + ) + + def test_created_llm_uses_new_key_on_next_request(self): + state = {"key": "old-key"} + llm = LLM(self._caps(), credential_resolver=lambda _env: state["key"]) + first = llm._build_kwargs([], None, None, None) + state["key"] = "new-key" + second = llm._build_kwargs([], None, None, None) + self.assertEqual(first["api_key"], "old-key") + self.assertEqual(second["api_key"], "new-key") + + def test_built_stream_request_keeps_key_snapshot(self): + state = {"key": "old-key"} + calls = [] + llm = LLM(self._caps(), credential_resolver=lambda _env: state["key"]) + + def completion(**kwargs): + calls.append(kwargs["api_key"]) + state["key"] = "new-key" + return iter([{"chunk": 1}, {"chunk": 2}]) + + with patch("core.llm.litellm.completion", side_effect=completion): + self.assertEqual(len(list(llm.chat_stream([]))), 2) + self.assertEqual(calls, ["old-key"]) + + +class ProviderMutationTests(unittest.TestCase): + def test_missing_master_key_refuses_database_save(self): + with ( + patch("core.provider_credentials.service.crypto_configured", return_value=False), + self.assertRaisesRegex(ProviderCredentialError, "MASTER_KEY"), + ): + replace_credentials( + "deepseek", {"api_key": "candidate"}, expected_revision=0, + updated_by=uuid4(), + ) + + def test_failed_candidate_never_opens_database_transaction(self): + with ( + patch("core.provider_credentials.service.crypto_configured", return_value=True), + patch("core.provider_credentials.service.test_provider", + return_value=TestResult("auth_error", "认证失败")), + patch("core.provider_credentials.service.session_scope") as scope, + self.assertRaises(ProviderCredentialError), + ): + replace_credentials( + "deepseek", {"api_key": "bad-key"}, expected_revision=0, + updated_by=uuid4(), + ) + scope.assert_not_called() + + def test_revision_conflict_does_not_overwrite(self): + class Session: + def execute(self, _statement): + return SimpleNamespace(rowcount=0) + + @contextmanager + def scope(): + yield Session() + + with ( + patch("core.provider_credentials.service.crypto_configured", return_value=True), + patch("core.provider_credentials.service.test_provider", + return_value=TestResult("normal", "正常")), + patch("core.provider_credentials.service.encrypt_secret", return_value="cipher"), + patch("core.provider_credentials.service.session_scope", scope), + self.assertRaises(RevisionConflict), + ): + replace_credentials( + "deepseek", {"api_key": "new-key"}, expected_revision=7, + updated_by=uuid4(), + ) + + def test_alert_cooldown_and_recovery(self): + now = datetime.now(timezone.utc) + low = TestResult("low_balance", "余额低") + self.assertEqual(_notification_transition("normal", None, low, now), "low_balance") + self.assertIsNone(_notification_transition("low_balance", now, low, now)) + self.assertEqual( + _notification_transition("low_balance", now - timedelta(days=1), low, now), + "low_balance", + ) + self.assertEqual( + _notification_transition("auth_error", now, TestResult("normal", "正常"), now), + "recovered", + ) + + def test_admin_payload_contains_no_ciphertext(self): + provider = BY_ID["deepseek"] + row = SimpleNamespace( + credentials={"api_key": "ciphertext-secret"}, + credential_hint={"api_key": "***1234"}, revision=2, + test_status="normal", test_detail="正常", last_tested_at=None, + balance_amount=None, balance_currency=None, + ) + with patch("core.provider_credentials.service.resolve_credentials") as resolve: + resolve.return_value = SimpleNamespace(source="database", values={"api_key": "plain-secret"}) + payload = _row_payload(provider, row) + encoded = json.dumps(payload, ensure_ascii=False) + self.assertNotIn("ciphertext-secret", encoded) + self.assertNotIn("plain-secret", encoded) + self.assertIn("***1234", encoded) + + +class ProviderLeaderTests(unittest.TestCase): + def test_unclaimed_round_does_not_test(self): + connection = Mock() + connection.exec_driver_sql.return_value.scalar.return_value = False + engine = MagicMock() + engine.connect.return_value.__enter__.return_value = connection + with ( + patch("core.provider_credentials.monitor.get_engine", return_value=engine), + patch("core.provider_credentials.monitor.due_provider_ids") as due, + ): + from core.provider_credentials.monitor import run_due_checks + self.assertEqual(run_due_checks(), 0) + due.assert_not_called() + + +class ProviderMigrationTests(unittest.TestCase): + def test_0039_creates_single_encrypted_credentials_table(self): + path = Path(__file__).resolve().parents[1] / "db" / "migrations" / "versions" / "20260902_1000_0039_provider_credentials.py" + spec = util.spec_from_file_location("migration_0039_test", path) + module = util.module_from_spec(spec) + assert spec.loader is not None + spec.loader.exec_module(module) + captured = {} + + def create_table(name, *items): + captured["name"] = name + captured["columns"] = {item.name for item in items if hasattr(item, "name")} + + with patch.object(module.op, "create_table", side_effect=create_table): + module.upgrade() + self.assertEqual(captured["name"], "provider_credentials") + self.assertTrue({"provider_id", "credentials", "credential_hint", "revision"} <= captured["columns"]) + self.assertNotIn("api_key", captured["columns"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_web_routes_nodb.py b/tests/test_web_routes_nodb.py index 6be0709..e131e76 100644 --- a/tests/test_web_routes_nodb.py +++ b/tests/test_web_routes_nodb.py @@ -117,6 +117,10 @@ class AuthGateTests(unittest.TestCase): ("POST", "/v1/tasks"), ("POST", "/v1/asr/transcribe"), ("GET", "/v1/admin/overview"), + ("GET", "/v1/admin/provider-credentials"), + ("PUT", "/v1/admin/provider-credentials/deepseek"), + ("POST", "/v1/admin/provider-credentials/deepseek/test"), + ("DELETE", "/v1/admin/provider-credentials/deepseek"), ("GET", "/v1/admin/sandbox/capacity"), ("GET", "/v1/admin/sandbox/packages"), ("GET", "/v1/admin/software-nodes"), diff --git a/web/admin.py b/web/admin.py index 87ae03b..7ad7fb1 100644 --- a/web/admin.py +++ b/web/admin.py @@ -237,6 +237,15 @@ class SetPlanRequest(BaseModel): plan: str = "" # 档位名(config/agent.yaml model_tiers 的 key);空串 = 清空 → 落 default 档 +class ProviderCredentialRequest(BaseModel): + credentials: dict[str, str] + expected_revision: int = Field(ge=0) + + +class ProviderCredentialDeleteRequest(BaseModel): + expected_revision: int = Field(ge=1) + + class ExternalSystemDefinitionRequest(BaseModel): provider: str = "generic_openapi" name: str @@ -312,6 +321,72 @@ def register_admin_routes(app: FastAPI, require_admin) -> None: "usage": usage_report.usage_overview(s, cutoff_7d), } + @app.get("/v1/admin/provider-credentials", tags=["admin"]) + def admin_provider_credentials(user_id: UUID = Depends(require_admin)): + from core.provider_credentials.service import list_providers + return {"results": list_providers()} + + @app.put("/v1/admin/provider-credentials/{provider_id}", tags=["admin"]) + def admin_replace_provider_credentials( + provider_id: str, + body: ProviderCredentialRequest, + user_id: UUID = Depends(require_admin), + ): + from core.provider_credentials.monitor import send_notification + from core.provider_credentials.service import ( + ProviderCredentialError, RevisionConflict, replace_credentials, + ) + try: + return replace_credentials( + provider_id, body.credentials, + expected_revision=body.expected_revision, + updated_by=user_id, notify=send_notification, + ) + except ValueError as exc: + raise HTTPException(404, "unknown provider") from exc + except RevisionConflict as exc: + raise HTTPException(409, str(exc)) from exc + except ProviderCredentialError as exc: + raise HTTPException(400, str(exc)) from exc + + @app.post("/v1/admin/provider-credentials/{provider_id}/test", tags=["admin"]) + def admin_test_provider_credentials( + provider_id: str, user_id: UUID = Depends(require_admin), + ): + from core.provider_credentials.monitor import send_notification + from core.provider_credentials.service import ( + ProviderCredentialError, test_current_credentials, + ) + try: + result = test_current_credentials(provider_id, notify=send_notification) + return { + "status": result.status, "detail": result.detail, + "balance": ( + {"amount": str(result.balance_amount), + "currency": result.balance_currency} + if result.balance_amount is not None else None + ), + } + except ValueError as exc: + raise HTTPException(404, "unknown provider") from exc + except ProviderCredentialError as exc: + raise HTTPException(400, str(exc)) from exc + + @app.delete("/v1/admin/provider-credentials/{provider_id}", tags=["admin"]) + def admin_delete_provider_credentials( + provider_id: str, + body: ProviderCredentialDeleteRequest, + user_id: UUID = Depends(require_admin), + ): + from core.provider_credentials.service import RevisionConflict, delete_override + try: + delete_override(provider_id, expected_revision=body.expected_revision) + except ValueError as exc: + raise HTTPException(404, "unknown provider") from exc + except RevisionConflict as exc: + raise HTTPException(409, str(exc)) from exc + return {"deleted": True} + @app.get("/v1/admin/sandbox/capacity", tags=["admin"]) def admin_sandbox_capacity(user_id: UUID = Depends(require_admin)): """宿主共享实时容量;只读文件/Docker 状态,不写 DB。""" diff --git a/web/app.py b/web/app.py index 1c0270e..aa74068 100644 --- a/web/app.py +++ b/web/app.py @@ -43,6 +43,7 @@ from .background import ( reap_stale_runs, start_disk_scanner, start_proc_sweeper, + start_provider_scanner, start_stats_logger, start_toolfail_scanner, ) @@ -124,6 +125,7 @@ def create_app() -> FastAPI: disk_scanner_task = start_disk_scanner(_cfg) stats_logger_task = start_stats_logger(app, run_max_workers) toolfail_task = start_toolfail_scanner() + provider_task = start_provider_scanner() scheduler_task = start_scheduler(app, _cfg) wechat_task, wechat_stop = start_wechat_inbound(app) sandbox_reaper_task = init_sandbox(app, _cfg) @@ -141,6 +143,7 @@ def create_app() -> FastAPI: await cancel_and_wait(disk_scanner_task) await cancel_and_wait(stats_logger_task) await cancel_and_wait(toolfail_task) + await cancel_and_wait(provider_task) await cancel_and_wait(scheduler_task) if wechat_task is not None: wechat_stop.set() diff --git a/web/background.py b/web/background.py index 539a6db..2a322c4 100644 --- a/web/background.py +++ b/web/background.py @@ -183,6 +183,26 @@ def start_toolfail_scanner() -> Optional[asyncio.Task]: return asyncio.create_task(_toolfail_scanner(), name="toolfail-scanner") +def start_provider_scanner() -> asyncio.Task: + """每分钟寻找到期 Provider;实际一轮由 PostgreSQL advisory lock 单实例执行。""" + async def _scanner() -> None: + from core.provider_credentials.monitor import run_due_checks + loop = asyncio.get_running_loop() + while True: + try: + await asyncio.sleep(5) + checked = await loop.run_in_executor(None, run_due_checks) + if checked: + print(f"[provider] checked {checked} provider(s)") + await asyncio.sleep(55) + except asyncio.CancelledError: + raise + except Exception as exc: + print(f"[provider] scanner error: {type(exc).__name__}") + await asyncio.sleep(60) + return asyncio.create_task(_scanner(), name="provider-scanner") + + def init_sandbox(app, cfg: dict) -> Optional[asyncio.Task]: """Sandbox pool(§7.5):仅当 ZCBOT_SANDBOX_BACKEND=docker 时启用。 diff --git a/web/static/admin.html b/web/static/admin.html index b3807c6..21eea1c 100644 --- a/web/static/admin.html +++ b/web/static/admin.html @@ -235,6 +235,11 @@ display: none; position: fixed; inset: 0; z-index: 140; justify-content: flex-end; background: rgba(18,24,30,.28); backdrop-filter: blur(1px); } + #s-provider-credentials button { + font-size: 12px; padding: 4px 8px; border: 1px solid var(--border); + border-radius: var(--r-md); background: #fff; cursor: pointer; + } + #s-provider-credentials button:disabled { opacity: .45; cursor: default; } .capacity-drawer.show { display: flex; } .capacity-drawer-panel { width: min(460px, calc(100vw - 32px)); height: 100%; display: flex; flex-direction: column; diff --git a/web/static/js/admin.js b/web/static/js/admin.js index 9e3fc9b..3c32f9b 100644 --- a/web/static/js/admin.js +++ b/web/static/js/admin.js @@ -16,6 +16,7 @@ const SECTIONS = [ ["s-users", "各用户"], ["s-storage", "存储"], ["s-sandbox-packages", "容器依赖"], ["s-windows-node", "专业软件节点"], + ["s-provider-credentials", "API 凭据"], ["s-external", "外部系统"], ["s-toolfail", "工具失败"], ]; @@ -54,6 +55,7 @@ let storagePage = 0; let packageRange = "30d"; let tiersData = null; // {tiers, default_tier, catalog};加载一次(改档位 / 看图例用) let externalDefinitions = []; +let providerCredentials = []; let externalUsers = []; let externalDefinitionsLoaded = false; let externalEditingId = ""; @@ -1154,6 +1156,87 @@ function renderMetrics(d) { renderOpsSummary(); } +function providerStatusLabel(status) { + return ({ normal: "正常", low_balance: "余额偏低", exhausted: "已耗尽", + auth_error: "认证失败", unreachable: "不可达", untested: "未测试" })[status] || status; +} + +function renderProviderCredentials() { + const rows = providerCredentials.map(row => { + const fields = (row.fields || []).map(f => `${escapeHtml(f.label)} ${escapeHtml(f.hint || "未配置")}`).join(" · "); + const balance = row.balance + ? `${escapeHtml(row.balance.currency || "")} ${escapeHtml(row.balance.amount || "")}` : "—"; + const tone = ["low_balance", "exhausted", "auth_error"].includes(row.test_status) ? "err" : + (row.test_status === "normal" ? "ok" : ""); + return `${escapeHtml(row.display_name)}` + + `
${escapeHtml(row.category)}` + + `${escapeHtml(row.source)}
${fields}` + + `${escapeHtml(providerStatusLabel(row.test_status))}` + + `
${escapeHtml(row.test_detail || "—")}` + + `${balance}
${row.last_tested_at ? fmtTimeAgo(row.last_tested_at) : "未测试"}` + + ` ` + + ``; + }).join("") || `暂无 Provider`; + $("s-provider-credentials").innerHTML = `

API 凭据

` + + `数据库密文优先;删除覆盖后回退环境变量。测试候选成功后才原子替换。
` + + `
` + + `${rows}
Provider来源与尾号状态余额 / 最近测试操作
`; + $("s-provider-credentials").onclick = async event => { + const tr = event.target.closest("tr[data-provider-id]"); + if (!tr) return; + const row = providerCredentials.find(item => item.provider_id === tr.dataset.providerId); + if (!row) return; + if (event.target.closest("[data-provider-replace]")) await replaceProviderCredentials(row); + else if (event.target.closest("[data-provider-test]")) await testProviderCredentials(row); + else if (event.target.closest("[data-provider-delete]")) await deleteProviderCredentials(row); + }; +} + +async function replaceProviderCredentials(row) { + if (row.billable_test && !confirm(`${row.display_name} 没有免费鉴权端点,本次测试会发起一次最小真实调用。继续?`)) return; + const credentials = {}; + for (const field of row.fields || []) { + const value = await dialogPrompt({ + title: `录入 ${row.display_name}`, + label: `${field.label}(当前 ${field.hint || "未配置"};不会回显原值)`, + placeholder: `输入新的 ${field.label}`, + value: "", maxLength: 4096, okText: "下一步", + }); + if (value === null) return; + if (!value.trim()) { message(`${field.label} 不可为空`, "error"); return; } + credentials[field.name] = value.trim(); + } + if (!await dialogConfirm({ title: `更换 ${row.display_name} 凭据?`, + message: "系统会先在内存中测试候选凭据;仅测试成功后才替换当前密文。", + okText: "测试并启用", danger: false })) return; + try { + const result = await apiSend("PUT", `/v1/admin/provider-credentials/${row.provider_id}`, + { credentials, expected_revision: row.revision || 0 }); + message(result.test_detail || "凭据已启用", result.test_status === "low_balance" ? "info" : "success"); + await loadProviderCredentials(); + } catch (err) { message("保存失败:" + (err.message || String(err)), "error", 6000); } +} + +async function testProviderCredentials(row) { + if (row.billable_test && !confirm(`${row.display_name} 测试会发起一次最小真实调用。继续?`)) return; + try { + const result = await apiSend("POST", `/v1/admin/provider-credentials/${row.provider_id}/test`, {}); + message(result.detail || providerStatusLabel(result.status), result.status === "normal" ? "success" : "info", 5000); + await loadProviderCredentials(); + } catch (err) { message("测试失败:" + (err.message || String(err)), "error", 6000); } +} + +async function deleteProviderCredentials(row) { + if (!await dialogConfirm({ title: `删除 ${row.display_name} 数据库覆盖?`, + message: "删除后立即回退到部署环境变量;若环境变量也未配置,该能力将不可用。", + okText: "删除覆盖", danger: true })) return; + try { + await apiSend("DELETE", `/v1/admin/provider-credentials/${row.provider_id}`, + { expected_revision: row.revision }); + await loadProviderCredentials(); + } catch (err) { message("删除失败:" + (err.message || String(err)), "error", 6000); } +} + function renderSandboxPackages(d) { const rows = d.rows || []; const opts = RANGE_OPTS.map(([v,l]) => ``).join(""); @@ -1315,6 +1398,14 @@ async function loadSoftwareNodes() { } catch (e) { /* overview 统一处理鉴权 */ } } +async function loadProviderCredentials() { + try { + const result = await apiGet("/v1/admin/provider-credentials"); + providerCredentials = result.results || []; + renderProviderCredentials(); + } catch (e) { /* overview 统一处理鉴权 */ } +} + async function loadSandboxCapacity() { try { sandboxCapacityData = await apiGet("/v1/admin/sandbox/capacity"); @@ -1336,6 +1427,7 @@ async function refresh() { loadStorage(storagePage); if (storagePage !== 0) loadStorageSummary(); loadSoftwareNodes(); + loadProviderCredentials(); loadExternalDefinitions(); loadToolFailures(); loadSandboxCapacity();