feat(admin): manage provider API credentials
This commit is contained in:
parent
aa8ccec738
commit
11108f31c2
|
|
@ -8,6 +8,10 @@
|
|||
|
||||
## Unreleased
|
||||
|
||||
- 单用户可同时运行的重型任务由 2 个提升到 3 个;任务等待执行容量时,对话会直接说明是当前用户、整机或宿主内存限制,获得槽位后自动继续。
|
||||
|
||||
- 管理员现在可在管理后台安全录入、测试和更换模型、媒体、检索及语音服务凭据,并查看来源、脱敏尾号和可用状态;数据库凭据可随时删除并回退原有环境配置。DeepSeek 余额低于 30 元、额度耗尽或认证失败时会主动提醒。
|
||||
|
||||
- Mermaid 图表采用更清晰的科研配色与更精致的图框,新生成的流程图还会按数据、处理、判断、结果等角色使用协调的语义色,减少单调的灰白图。
|
||||
|
||||
- 管理后台总览按“当前运行”和“运营资源”重新组织,执行容量改为精简摘要并可从右侧详情面板查看队列、容器、用户占用与宿主资源,异常状态和近期指标也更容易识别。
|
||||
|
|
|
|||
14
DESIGN.md
14
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)
|
||||
|
|
|
|||
|
|
@ -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 后台协程按域析出
|
||||
|
|
|
|||
13
RUN.md
13
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 = "<dedicated-eval-user-jwt>"
|
|||
| `GET /v1/skills/{name}` | 返某 skill 完整 SKILL.md 正文(前端「技能」modal 点开查看);同名按 user wins | 必填 |
|
||||
| `DELETE /v1/skills/{name}` | 删当前 user 私有 skill(`.skills/<name>/` 整目录);只删 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 = "<dedicated-eval-user-jwt>"
|
|||
| `GET /v1/models` | 列 chat LLM 模型清单(扫 `config/models/*.yaml`),前端顶栏切换 / 新建对话框下拉用 | 必填 |
|
||||
| `GET /v1/image_models` | 列图像生成 variant 清单(扫 `config/media/doubao.yaml` image 段),前端"生图"下拉用;yaml 无 image variant → 空列表 → UI 隐藏下拉 | 必填 |
|
||||
|
||||
**SSE 事件**(每帧 `event: <type>` + `data: <JSON>`):建连时若当前 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: <type>` + `data: <JSON>`):建连时若当前 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 规则。生产部署必填。
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
59
core/llm.py
59
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):
|
||||
|
|
|
|||
|
|
@ -0,0 +1,5 @@
|
|||
"""Admin 管理的平台 Provider 凭据控制面。"""
|
||||
|
||||
from .runtime import resolve_credentials, resolve_secret
|
||||
|
||||
__all__ = ["resolve_credentials", "resolve_secret"]
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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"),
|
||||
|
|
|
|||
75
web/admin.py
75
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。"""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 时启用。
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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 `<tr data-provider-id="${escapeHtml(row.provider_id)}"><td>${escapeHtml(row.display_name)}`
|
||||
+ `<br/><span class="muted">${escapeHtml(row.category)}</span></td>`
|
||||
+ `<td><span class="chip">${escapeHtml(row.source)}</span><br/><span class="muted">${fields}</span></td>`
|
||||
+ `<td><span class="chip ${tone}">${escapeHtml(providerStatusLabel(row.test_status))}</span>`
|
||||
+ `<br/><span class="muted" title="${escapeHtml(row.test_detail || "")}">${escapeHtml(row.test_detail || "—")}</span></td>`
|
||||
+ `<td>${balance}<br/><span class="muted">${row.last_tested_at ? fmtTimeAgo(row.last_tested_at) : "未测试"}</span></td>`
|
||||
+ `<td><button data-provider-replace>录入/更换</button> <button data-provider-test ${row.configured ? "" : "disabled"}>测试</button> `
|
||||
+ `<button data-provider-delete ${row.source === "database" ? "" : "disabled"}>删除覆盖</button></td></tr>`;
|
||||
}).join("") || `<tr><td colspan="5" class="empty">暂无 Provider</td></tr>`;
|
||||
$("s-provider-credentials").innerHTML = `<div class="card"><div class="card-head"><h2>API 凭据</h2>`
|
||||
+ `<span class="sublabel">数据库密文优先;删除覆盖后回退环境变量。测试候选成功后才原子替换。</span></div>`
|
||||
+ `<div class="scroll-x"><table><thead><tr><th>Provider</th><th>来源与尾号</th><th>状态</th><th>余额 / 最近测试</th><th>操作</th></tr></thead>`
|
||||
+ `<tbody>${rows}</tbody></table></div></div>`;
|
||||
$("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]) => `<option value="${v}" ${v===packageRange?"selected":""}>${l}</option>`).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();
|
||||
|
|
|
|||
Loading…
Reference in New Issue