diff --git a/.gitignore b/.gitignore index 131c027..0bfbbea 100644 --- a/.gitignore +++ b/.gitignore @@ -26,3 +26,7 @@ cmd/gui/dist/ # local design/working notes (not part of the shipped repo) /plan.md + +# scratch build/probe dirs used while debugging (GOTMPDIR, throwaway binaries) +.build-work/ +.probe/ diff --git a/cmd/gui/main.js b/cmd/gui/main.js index ab33943..abc6ea9 100644 --- a/cmd/gui/main.js +++ b/cmd/gui/main.js @@ -671,6 +671,78 @@ ipcMain.handle( (e) => !!BrowserWindow.fromWebContents(e.sender)?.isMaximized(), ); +// ---- plugin management over IPC ------------------------------------------- +// +// The renderer cannot call the embedded core directly: it has no key and no +// network identity, and the core binds a loopback port that only the main +// process knows about. So every plugin action is proxied through the main +// process, which already knows how to obtain the admin key (unsealViaCore). +// +// The proxy is deliberately a raw (method, path, body) pass-through rather than +// a fixed set of commands. A fixed set would have to be extended for every new +// plugin endpoint, and the one thing worse than "no button for this" is "a +// button that silently does nothing" — with a pass-through the renderer can talk +// to any /api/plugins route the core grows, and the path is validated here so +// this channel cannot be used to reach arbitrary endpoints. +function pluginProxy(req) { + const { method, path, body } = req || {}; + const M = ["GET", "POST", "PUT", "DELETE"]; + if (!M.includes(method)) throw new Error("bad method: " + method); + // The path must stay inside the plugin namespace. A prefix check alone would + // still allow /api/plugins/../keys, so reject any traversal outright. + if (typeof path !== "string" || !path.startsWith("/api/plugins")) { + throw new Error("path must start with /api/plugins"); + } + if (path.includes("..") || path.includes("\\")) { + throw new Error("path traversal rejected"); + } + const key = unsealViaCore(); + if (!key) throw new Error("no admin key available yet"); + return new Promise((resolve, reject) => { + const u = new URL(embeddedBaseUrl() + path); + const data = body == null ? null : JSON.stringify(body); + const headers = { Authorization: "Bearer " + key }; + if (data) { + headers["Content-Type"] = "application/json"; + headers["Content-Length"] = Buffer.byteLength(data); + } + const r = http.request( + { + hostname: u.hostname, + port: u.port, + path: u.pathname + u.search, + method, + headers, + }, + (res) => { + let raw = ""; + res.setEncoding("utf8"); + res.on("data", (c) => (raw += c)); + res.on("end", () => { + let parsed = null; + try { + parsed = raw ? JSON.parse(raw) : null; + } catch (e) { + parsed = { raw }; + } + if (res.statusCode >= 400) { + const msg = + (parsed && parsed.error && parsed.error.message) || + "HTTP " + res.statusCode; + reject(new Error(msg)); + return; + } + resolve(parsed); + }); + }, + ); + r.on("error", reject); + if (data) r.write(data); + r.end(); + }); +} + +ipcMain.handle("plugins:proxy", (_e, req) => pluginProxy(req)); ipcMain.handle("core:state", () => ({ running: coreStarted() && coreReady, ready: coreReady, diff --git a/cmd/gui/package-lock.json b/cmd/gui/package-lock.json index 96dd73e..21f79ad 100644 --- a/cmd/gui/package-lock.json +++ b/cmd/gui/package-lock.json @@ -582,7 +582,7 @@ } }, "node_modules/@peculiar/webcrypto": { - "version": "1.7.1", + "version": "1.7.5", "resolved": "https://registry.npmjs.org/@peculiar/webcrypto/-/webcrypto-1.7.1.tgz", "integrity": "sha512-ODOov0sGMJMf3jPonOkgGqPknTsu+DdQ7kD++gz8aI+aFMOMHFbWAA2taqXXVTdP+OTOQR/znGvSpmkeI0WTYQ==", "dev": true, @@ -3335,7 +3335,7 @@ } }, "node_modules/resedit": { - "version": "1.7.2", + "version": "1.7.5", "resolved": "https://registry.npmjs.org/resedit/-/resedit-1.7.2.tgz", "integrity": "sha512-vHjcY2MlAITJhC0eRD/Vv8Vlgmu9Sd3LX9zZvtGzU5ZImdTN3+d6e/4mnTyV8vEbyf1sgNIrWxhWlrys52OkEA==", "dev": true, diff --git a/cmd/gui/preload.js b/cmd/gui/preload.js index c3cfca8..f9c46e4 100644 --- a/cmd/gui/preload.js +++ b/cmd/gui/preload.js @@ -16,6 +16,13 @@ contextBridge.exposeInMainWorld("modelrouter", { key: () => ipcRenderer.invoke("core:key"), onState: (cb) => ipcRenderer.on("core:state", (_e, d) => cb(d)), }, + plugins: { + // Raw pass-through to the embedded core's /api/plugins surface. The main + // process validates the path and attaches the admin key; the renderer never + // sees either. + request: (method, path, body) => + ipcRenderer.invoke("plugins:proxy", { method, path, body }), + }, settings: { get: () => ipcRenderer.invoke("settings:get"), set: (patch) => ipcRenderer.invoke("settings:set", patch), diff --git a/cmd/gui/renderer/app.js b/cmd/gui/renderer/app.js index 2d8d2ce..553a30b 100644 --- a/cmd/gui/renderer/app.js +++ b/cmd/gui/renderer/app.js @@ -155,6 +155,7 @@ async function openSettings() { $("#set-tray").checked = !!state.settings.minimizeToTray; $("#settings-overlay").style.display = "flex"; renderRail(); + loadPlugins(); } function closeSettings() { $("#settings-overlay").style.display = "none"; @@ -182,6 +183,122 @@ async function saveSettings() { } } +// esc / escAttr escape text for innerHTML. The WebUI has its own copies; the +// shell needs its own because renderer/app.js is a separate document that +// never loads index.html's script. +function esc(s) { + return String(s == null ? "" : s).replace( + /[&<>"']/g, + (c) => ({ "&": "&", "<": "<", ">": ">", '"': """, "'": "'" })[c], + ); +} +function escAttr(s) { + return esc(s).replace(/`/g, "`"); +} + +// ===== plugin management ===== +// +// The desktop shell manages plugins through the embedded core's /api/plugins +// surface, proxied over IPC (see plugins:proxy in the main process). The +// renderer never holds the admin key. +// +// Scope note: the desktop build has no plugin_dir configured by default, so this +// panel normally reports "plugins disabled" with the one-line fix. That is a +// deliberate state, not an error — the packaged profile is a per-user directory +// and seeding a plugin tree into someone's home without asking would be rude. + +async function loadPlugins() { + const list = $("#pl-list"); + const hint = $("#set-plugins-hint"); + if (!list || !hint) return; + let j; + try { + j = await window.modelrouter.plugins.request("GET", "/api/plugins"); + } catch (e) { + hint.textContent = "内核未就绪:" + e.message; + list.innerHTML = ""; + return; + } + if (!j.plugin_dir) { + hint.innerHTML = + '未配置 plugin_dir,插件功能未启用。在 config.yaml 加一行后重启内核即可。'; + list.innerHTML = ""; + return; + } + const rows = j.on_disk || []; + const active = rows.filter((p) => p.loaded && !p.disabled).length; + const broken = rows.filter((p) => !p.loaded).length; + hint.textContent = + `${rows.length} 个插件 · ${active} 个启用中` + + (broken ? ` · ${broken} 个加载失败` : ""); + list.innerHTML = rows.length + ? rows + .map((p) => { + const cls = !p.loaded ? "pl-broken" : p.disabled ? "pl-off" : "pl-on"; + const label = !p.loaded + ? "加载失败" + : p.disabled + ? "已禁用" + : "启用中"; + const btn = p.loaded + ? `` + : ""; + const builtin = p.builtin + ? '内置' + : ""; + return `
+
${esc(p.name)}${builtin}${label}
+ ${p.description ? `
${esc(p.description)}
` : ""} + ${p.error ? `
${esc(String(p.error).slice(0, 160))}
` : ""} +
${btn}
+
`; + }) + .join("") + : '
插件目录为空
'; + list.querySelectorAll('button[data-act="toggle"]').forEach((b) => { + b.onclick = () => togglePlugin(b.dataset.name, b.dataset.en === "1"); + }); +} + +async function togglePlugin(name, disabled) { + try { + await window.modelrouter.plugins.request("PUT", `/api/plugins/${encodeURIComponent(name)}`, { + enabled: disabled, + }); + toast(disabled ? `已禁用 ${name}` : `已启用 ${name}`); + await loadPlugins(); + } catch (e) { + toast(e.message, true); + } +} + +async function enableAllPlugins() { + let j; + try { + j = await window.modelrouter.plugins.request("GET", "/api/plugins"); + } catch (e) { + return toast(e.message, true); + } + const off = (j.on_disk || []).filter((p) => p.loaded && p.disabled); + for (const p of off) { + try { + await window.modelrouter.plugins.request( + "PUT", + `/api/plugins/${encodeURIComponent(p.name)}`, + { enabled: true }, + ); + } catch (e) { + toast(`${p.name}: ${e.message}`, true); + } + } + toast(off.length ? `已启用 ${off.length} 个插件` : "没有处于禁用状态的插件"); + await loadPlugins(); +} + // ===== theme ===== function applyTheme() { document.documentElement.dataset.theme = state.theme; @@ -197,6 +314,10 @@ function init() { $("#tb-close").onclick = () => window.modelrouter.win.close(); $("#tb-settings").onclick = openSettings; $("#rail-settings").onclick = openSettings; + const plReload = document.getElementById("pl-reload"); + if (plReload) plReload.onclick = loadPlugins; + const plAll = document.getElementById("pl-toggle-all"); + if (plAll) plAll.onclick = enableAllPlugins; $("#rail-autostart").onclick = toggleAutoStart; $("#rail-silent").onclick = toggleSilent; $("#rail-theme").onclick = () => { diff --git a/cmd/gui/renderer/index.html b/cmd/gui/renderer/index.html index 5e16ec6..86634c0 100644 --- a/cmd/gui/renderer/index.html +++ b/cmd/gui/renderer/index.html @@ -207,6 +207,15 @@ > 关闭时最小化到托盘点关闭按钮隐藏到系统托盘 +
+ + 加载中… +
+
+
+ + +
diff --git a/cmd/gui/renderer/style.css b/cmd/gui/renderer/style.css index b3661aa..05dd59c 100644 --- a/cmd/gui/renderer/style.css +++ b/cmd/gui/renderer/style.css @@ -596,3 +596,81 @@ html[data-theme="dark"] .overlay { #toast.err { border-color: var(--danger); } + +/* ===== plugin management panel ========================================= + * The existing .ghost/.primary rules are scoped to `.form .actions`, so a + * button outside that selector gets browser defaults. The plugin rows live in + * their own list, hence their own rules — reusing a scoped class here would have + * produced unstyled buttons that still worked, which is the kind of thing that + * looks fine until someone themes the shell. + */ +.pl-list { + display: flex; + flex-direction: column; + gap: 8px; + margin: 8px 0 4px; +} +.pl-item { + border: 1px solid var(--line); + border-radius: 9px; + padding: 10px 12px; +} +.pl-item.pl-broken { + border-color: var(--danger, #d1435b); +} +.pl-head { + display: flex; + align-items: center; + gap: 8px; + font-size: 13px; +} +.pl-builtin { + font-size: 10px; + padding: 1px 6px; + border-radius: 999px; + background: var(--primary-50); + color: var(--primary-h); +} +.pl-state { + margin-left: auto; + font-size: 11px; + color: var(--muted); +} +.pl-item.pl-broken .pl-state { + color: var(--danger, #d1435b); +} +.pl-desc { + font-size: 12px; + color: var(--muted); + margin-top: 3px; +} +.pl-err { + font-size: 11px; + color: var(--danger, #d1435b); + margin-top: 4px; + word-break: break-word; +} +.pl-acts { + margin-top: 8px; + display: flex; + gap: 8px; +} +.pl-acts button { + padding: 5px 12px; + font-size: 12px; + border-radius: 7px; + border: 1px solid var(--line); + background: var(--bg-s2, #fff); + color: var(--fg, inherit); + cursor: pointer; + transition: all 0.15s; +} +.pl-acts button:hover { + border-color: var(--primary); + color: var(--primary-h); +} +.pl-empty { + font-size: 12px; + color: var(--muted); + padding: 10px 0; +} diff --git a/deploy.sh b/deploy.sh index a4369ad..27712fe 100755 --- a/deploy.sh +++ b/deploy.sh @@ -301,6 +301,60 @@ verify_master_key() { return 1 } +# ---------- 适配器备份保留策略 ---------- +# 每次部署都新建一个 adapters.bak.<时间戳> 目录,而回滚只用到最近一次 +# ($BACKUP_BIN / $BACKUP_CONFIG 都是单文件覆盖)。多出来的目录从没人清理, +# 实测线上已积累 71 个,/etc/llmsproxy 因此涨到 122M。 +# +# 保留最近 KEEP_ADAPTER_BACKUPS 份足够回滚,同时给目录数设上限——否则 +# 一次误配置(比如 adapter_dir 指错)就可能在几秒内造出成百上千个目录。 +# 只删名字严格匹配 adapters.bak.<14位时间戳> 的目录,避免误伤人工放的目录。 +KEEP_ADAPTER_BACKUPS=5 + +prune_adapter_backups() { + local base="/etc/llmsproxy" + [[ -d "$base" ]] || return 0 + + # 先按数量上限硬裁:即使时间戳排序失效也不会无上限增长。 + local all + mapfile -t all < <(find "$base" -maxdepth 1 -type d -name 'adapters.bak.*' -printf '%f\n' | sort) + local cap=$((KEEP_ADAPTER_BACKUPS * 4)) + if (( ${#all[@]} > cap )); then + warn "适配器备份目录有 ${#all[@]} 个(异常),裁到 $cap" + local i=0 + for d in "${all[@]}"; do + i=$((i + 1)) + # 从最旧的开始删(sort 后升序)。名字不规范的跳过不删。 + if (( i <= ${#all[@]} - cap )); then + if [[ "$d" =~ ^adapters\.bak\.[0-9]{14}$ ]]; then + rm -rf "${base:?}/$d" + else + warn "跳过名字不规范的备份目录(不删): $d" + fi + fi + done + mapfile -t all < <(find "$base" -maxdepth 1 -type d -name 'adapters.bak.*' -printf '%f\n' | sort) + fi + + # 再按时间保留最近 KEEP_ADAPTER_BACKUPS 份(sort 后最新在末尾)。 + local total=${#all[@]} i=0 + for d in "${all[@]}"; do + i=$((i + 1)) + # i <= total-KEEP 的是较旧的,要删。 + if (( i <= total - KEEP_ADAPTER_BACKUPS )); then + if [[ "$d" =~ ^adapters\.bak\.[0-9]{14}$ ]]; then + rm -rf "${base:?}/$d" + else + warn "跳过名字不规范的备份目录(不删): $d" + fi + fi + done + + local left + left=$(find "$base" -maxdepth 1 -type d -name 'adapters.bak.*' | wc -l) + log " 适配器备份保留最近 $KEEP_ADAPTER_BACKUPS 份(当前剩 $left 个)" +} + # ---------- 同步适配器 ---------- sync_adapters() { log "同步适配器" @@ -318,6 +372,7 @@ sync_adapters() { cp -rf "$TARGET_ADAPTERS/"*.lua "$BACKUP_DIR/" 2>/dev/null || true log " 旧适配器已备份到 $BACKUP_DIR" fi + prune_adapter_backups cp -f "$SRC_ADAPTERS/"*.lua "$TARGET_ADAPTERS/" chmod 0644 "$TARGET_ADAPTERS/"*.lua diff --git a/docs/git-workflow-en.md b/docs/git-workflow-en.md index 8c3064b..0ffea85 100644 --- a/docs/git-workflow-en.md +++ b/docs/git-workflow-en.md @@ -15,7 +15,7 @@ deployable. |---|---|---|---| | Main | `main` | permanent | Only long-lived branch. Always deployable. Accumulates the next version. | | Feature | `feature/` | short (dev → merge → delete) | New features / ordinary fixes. Born from `main`, merged back into `main`. | -| Release | `release/vX.Y.Z` | one version cycle | Cut from `main`, tagged for release. Version-specific hotfixes land here. | +| Release | `release/vX.Y.x` | one version cycle | Cut from `main`, tagged for release. Version-specific hotfixes land here. | ## Change flow (important) @@ -26,15 +26,21 @@ deployable. │ │ │ cut │ cut ▼ ▼ - release/v1.4.2 release/v1.4.3 + release/v1.6.x release/v1.7.x │ │ - tag: v1.4.2 tag: v1.4.3 + tag: v1.6.0 tag: v1.7.6 │ │ hotfix ◄─────┘ hotfix ◄─────┘ │ │ └── cherry-pick back ──────────┘ ``` +The patch position in the branch name is a literal `x`, while the tag carries +the concrete version: one `release/v1.7.x` can hold tags v1.7.0 … v1.7.6. +Spanning several patches on a single minor branch is deliberate — patches are +revisions of the same feature batch, hotfixes land on one branch, and back-port +to main never has to resolve dependencies between several release branches. + ### Key rules 1. **main is always deployable**: never leave half-done work on `main`. @@ -43,9 +49,9 @@ deployable. then `git merge --no-ff feature/xxx` (or squash) when done. 3. **Release = cut a release branch from main + tag**: ```bash - git checkout -b release/v1.4.2 main - git tag -a v1.4.2 -m "ModelRouter v1.4.2" - git push origin release/v1.4.2 v1.4.2 + git checkout -b release/v1.7.x main + git tag -a v1.7.0 -m "ModelRouter v1.7.0" + git push origin release/v1.7.x v1.7.0 ``` Build installers and upload the GitCode Release from this tag so the published state is exactly reproducible. @@ -54,7 +60,7 @@ deployable. an already-released branch (unless you deliberately ship a minor revision). 5. **Hotfixes MUST flow back to main**: ```bash - git checkout release/v1.4.2 # fix in the release branch + git checkout release/v1.7.x # fix in the release branch git commit -m "fix: ..." git checkout main git cherry-pick # and into main @@ -67,12 +73,23 @@ deployable. When the next version ships, the previous release branch retires: - **Default: delete the remote release branch** - (`git push origin :release/v1.4.2`). All hotfixes were already + (`git push origin :release/v1.7.x`). All hotfixes were already cherry-picked into main, so main contains everything; no merge needed. - **Long-term maintenance** (e.g. an enterprise client pinned to an old version): keep the branch, accept only security fixes, keep the commit-then-cherry-pick loop. +> **Where practice diverged from this section (checked 2026-10-01)**: +> `release/v1.4.x` and `release/v1.5.x` still exist locally and on the remote, +> so "retire the previous branch when the next version ships" was never +> carried out. Keeping them is harmless (hotfixes were back-ported), but it +> contradicts the rule above and makes a reader wonder whether they should be +> there at all. **Feature branches, by contrast, are cleaned up**: +> `feature/key-quota-control`, `feature/toolcall-id-sanitize`, +> `feature/anthropic-usage-cache` and `feature/agentrouter-id-sanitize` were +> deleted on 2026-10-01 after confirming with a per-commit `git patch-id` +> comparison that their work had already landed in main. + ## Explicit non-goals - **Never rebase main**: main's history stays append-only; anyone pulling gets @@ -134,8 +151,8 @@ pain points: 2. No feature branches meant two independent efforts could not proceed in parallel without colliding. -With release branches: the published state = `release/vX.Y.Z` branch + -`vX.Y.Z` tag, exactly reproducible; hotfixes have a clear landing spot; main -stays "latest + all fixes + deployable". +With release branches: the published state = `release/vX.Y.x` branch + +the concrete `vX.Y.Z` tag, exactly reproducible; hotfixes have a clear landing +spot; main stays "latest + all fixes + deployable". > 中文版见 [docs/git-workflow.md](git-workflow.md)。 \ No newline at end of file diff --git a/docs/git-workflow.md b/docs/git-workflow.md index 363d5a7..af8722b 100644 --- a/docs/git-workflow.md +++ b/docs/git-workflow.md @@ -13,7 +13,7 @@ ModelRouter 采用 **GitHub Flow + 发布分支** 模型:`main` 是唯一长 |---|---|---|---| | 主分支 | `main` | 永久 | 唯一长命分支。永远可部署。积攒下一个版本的功能。 | | 特性分支 | `feature/<描述>` | 短命(开发→合并即删) | 新特性 / 一般 bug 修复。从 `main` 开出,完成后合回 `main`。 | -| 发布分支 | `release/vX.Y.Z` | 一个版本周期 | 从 `main` 分出,打 tag 发布。该版本生命周期内的 hotfix 都提交在此分支。 | +| 发布分支 | `release/vX.Y.x` | 一个版本周期 | 从 `main` 分出,打 tag 发布。该版本生命周期内的 hotfix 都提交在此分支。 | ## 变更流向(重要) @@ -24,15 +24,20 @@ ModelRouter 采用 **GitHub Flow + 发布分支** 模型:`main` 是唯一长 │ │ │ 切出 │ 切出 ▼ ▼ - release/v1.4.2 release/v1.4.3 + release/v1.6.x release/v1.7.x │ │ - tag: v1.4.2 tag: v1.4.3 + tag: v1.6.0 tag: v1.7.6 │ │ hotfix ◄─────┘ hotfix ◄─────┘ │ │ └── cherry-pick 回 main ───────┘ ``` +分支名里 patch 位是**字面的 x**,而 tag 打具体版本号:`release/v1.7.x` 这一条 +发布分支上的 tag 可以有 v1.7.0 … v1.7.6 多个。一条 minor 分支跨多个 patch 是 +刻意的:patch 是同一批功能的不同修订,hotfix 落在同一条分支上,回流 main 时 +也不必处理多条 release 分支之间的依赖。 + ### 关键规则 1. **main 永远可部署**:不在 main 上留半成品。任何未完成的工作必须在特性分支上。 @@ -40,16 +45,16 @@ ModelRouter 采用 **GitHub Flow + 发布分支** 模型:`main` 是唯一长 开发完 `git merge --no-ff feature/xxx` 或 squash 合回。 3. **发布 = 从 main 切 release 分支 + 打 tag**: ```bash - git checkout -b release/v1.4.2 main - git tag -a v1.4.2 -m "ModelRouter v1.4.2" - git push origin release/v1.4.2 v1.4.2 + git checkout -b release/v1.7.x main + git tag -a v1.7.0 -m "ModelRouter v1.7.0" + git push origin release/v1.7.x v1.7.0 ``` 构建安装包、上传 GitCode Release 都基于这个 tag,保证可精确回溯发布态。 4. **版本生命周期内只收该版本的 hotfix**:新特性一律并入 `main` 等下一个版本, 绝不塞进已发布的 release 分支(除非主动选择在该版本内发次要版)。 5. **hotfix 必须回流 main**: ```bash - git checkout release/v1.4.2 # 在发布分支提交修复 + git checkout release/v1.7.x # 在发布分支提交修复 git commit -m "fix: ..." git checkout main git cherry-pick # 回主分支 @@ -61,11 +66,20 @@ ModelRouter 采用 **GitHub Flow + 发布分支** 模型:`main` 是唯一长 下一个版本发布时,上一个 release 分支退役: -- **默认:直接删除远端 release 分支**(`git push origin :release/v1.4.2`)。 +- **默认:直接删除远端 release 分支**(`git push origin :release/v1.7.x`)。 因为 hotfix 都已逐个 cherry-pick 回 main,main 已包含全部修复,无需再合并。 - **如需要长期维护旧版**(例如企业大客户卡在旧版本):保留分支,仅 stopship 接受 该版本的安全修复,继续走「提交 + cherry-pick 回 main」循环。 +> **实践与本节的历史出入(2026-10-01 核对)**:`release/v1.4.x` 与 `release/v1.5.x` +> 至今仍在本地与远端,说明"下一个版本发布就删上一个分支"实际没有执行。 +> 保留无害(hotfix 已回流),但它与上面写的规则不一致,读文档的人会以为 +> 这些分支不该存在。**特性分支则确实在清理**:`feature/key-quota-control`、 +> `feature/toolcall-id-sanitize`、`feature/anthropic-usage-cache`、 +> `feature/agentrouter-id-sanitize` 四个分支在 2026-10-01 删除——它们的工作 +> 早已全部进入 main(逐提交用 `git patch-id` 比对确认),留着只是给下个版本 +> 制造 cherry-pick/merge 陷阱。 + ## 明确不做的事 - **不 rebase main**:`main` 的历史保持追加式,任何人拉取后 `git pull` 都得到直接可用的历史。 diff --git a/docs/plugins.md b/docs/plugins.md new file mode 100644 index 0000000..a759602 --- /dev/null +++ b/docs/plugins.md @@ -0,0 +1,507 @@ +# 插件系统(Plugin System) + +> **English**: this document is the reference for writing ModelRouter plugins. +> The Chinese version is the primary one; section titles map 1:1. + +ModelRouter 的插件是**单个 `.lua` 文件**,放在 `config.yaml` 的 `plugin_dir` 目录里。 +插件能做两件事: + +1. **挂钩子**:在请求流水线的若干 stage 上注册回调,看到每个请求的完整信息, + 并可以把结果累加进自己的状态。 +2. **贡献界面**:在启动时返回 HTML / CSS / JS,由内核注入 WebUI——可以是一整个 + 新页面,也可以是往现有页面里追加一个组件。 + +两者互相独立:只想统计请求数的插件不必碰界面;只想加个仪表盘的插件不必碰钩子。 + +--- + +## 1. 快速上手 + +一个最小的插件: + +```lua +-- plugins/hello.lua +local plugin = { + name = "hello", + version = "1.0.0", + description = "示例插件", + author = "you", +} + +-- 声明钩子 +plugin.hooks = { + request_end = "on_request_end", +} + +-- 钩子实现 +function plugin.on_request_end(payload) + -- payload 是解码后的 table,不是 JSON 字符串 + log("info", string.format("%s via %s: %d prompt tokens", + payload.model, payload.source, payload.prompt_tokens or 0)) + return nil -- 最后一个 stage 没有下游,return 无意义 +end + +-- 贡献界面 +plugin.ui = { + page = { + page_id = "hello", -- kebab-case + title = "Hello", + icon = "👋", + order = 90, -- 侧栏排序 + mount = [[
hello
]], + }, +} + +return plugin -- 必须返回一个 table +``` + +放进 `plugin_dir` 后重启即生效。`GET /api/plugins` 确认它被加载了。 + +--- + +## 2. 加载与生命周期 + +``` +core.New + └─ lua.NewVM(adapter_dir).Start() 适配器状态 + └─ lua.NewPlugins(vm, plugin_dir) + ├─ SeedBundled() 仅当目录不存在时写入内置插件(目前是 billing) + └─ LoadDir() 按文件名字典序逐个加载 +``` + +**加载失败不影响网关启动。** 一个语法错误的插件会被记录在 `GET /api/plugins` +的 `error` 字段里,永远不会被调用。这与适配器一致,但理由更强:插件是可选的 +第三方扩展,因为一个 `.lua` 打错字就让网关起不来是错误的取舍。 + +**目录一旦存在就是权威的。** 与适配器同规则:首启会 seed 内置插件,之后目录里 +的文件说了算,删除或编辑内置插件都是真实生效的操作。 + +### 2.1 热更新 + +| 方式 | 效果 | +|---|---| +| `POST /api/plugins {name, code}` | 写文件 + 立即加载新版本(旧的 Lua 状态被关闭重建,**累计量清零**) | +| `DELETE /api/plugins/{name}` | 删文件 + 卸载 | +| 改文件后 `POST` 同名 | 同上 | + +改文件但**不** POST,需要重启才生效。 + +--- + +## 3. 流水线 stage + +一个请求依次经过三个 stage。插件可以为任意 stage 注册钩子;未注册的 stage +被忽略,所以插件不会因为网关将来新增 stage 而报错。 + +``` + 客户端请求 + │ + ┌─────────▼──────────┐ + │ request_start │ 已解析、已鉴权,尚未选源 + │ · type │ "chat" | "stream" | "image" + │ · model │ 客户端请求的原始 model("AUTO" 也在这里) + │ · key / role │ 掩码后的网关 key id("***a1b2c3")与角色 + │ · source │ 空(还没选源) + │ · stream │ + │ · messages_count │ + │ · tools_count │ + │ · ts │ unix 秒 + └─────────┬──────────┘ + │ (调度:tier 遍历 → 槽位轮转 → 冷却/配额过滤) + ┌─────────▼──────────┐ + │ routed │ 已选定 (source, model),尚未发往上游 + │ · source / model │ 实际选中的 + │ · tier │ AUTO 链的档位;直连 = -1;AUTO = -2 + │ · stream / key │ + └─────────┬──────────┘ + │ (HTTP 往返 / SSE 流) + ┌─────────▼──────────┐ + │ request_end │ 每个请求恰好一次,成功失败都触发 + │ · ok / status │ + │ · latency_ms │ + │ · first_byte_ms │ 流式的首字节时间 + │ · prompt_tokens │ 上游真实 usage,缺失时为字节估算 + │ · completion_tokens + │ · cache_hit_tokens / cache_miss_tokens + │ · image_count │ 生图数量(图片不计 token) + │ · error │ 失败原因,成功时为 "" + │ · time │ unix **毫秒** + └─────────┬──────────┘ + │ + 写审计 + 聚合统计 +``` + +### 3.1 触发点在哪 + +| stage | 代码位置 | 说明 | +|---|---|---| +| `request_start` | `gateway/chat.go` `handleChat` / `fireImageStart` | 每个请求一次(chat 与生图各一条),在配额闸门**之前** | +| `chain_step` | `gateway/chat.go` `chainTraceSink` | **仅 AUTO 路径**,每步一次 | +| `routed` | `singleChat` / `streamChat` / `singleChatAuto` / `streamChatAuto` | 成功选定源之后,各一次 | +| `request_end` | `gateway/chat.go` `writeRec` | **所有出口的唯一汇合点**,每个请求一次 | + +### 3.1b 为什么单独有 `chain_step` + +`routed` 只在**遍历结束后**触发一次,只带最终胜出的槽位。所以 +"tier 1 冷却所以降级到 tier 3"和"tier 1 正常接单"在它眼里**完全一样**—— +而这恰恰是优先级链存在的全部理由。 + +`chain_step` 补上这条信息,四种 `kind`: + +| kind | 含义 | 何时产生 | +|---|---|---| +| `tier_skip` | 整档被跳过 | 该档所有槽位冷却中/配额用尽 | +| `slot_fail` | 某个槽位硬失败 | 上游报错 / 适配器输出不可用 | +| `tier_busy` | 整档全忙且有界等待超时 | 2s 内没等到空位 | +| `selected` | 这个槽位接了单 | 每次成功遍历**恰好一次**,且是最后一步 | + +顺序保证:所有 `chain_step` 都在 `routed` 之前,`selected` 是最后一步。 +所以只订阅 `request_end` 的插件也能拿到轨迹摘要(见下)。 + +### 3.1c `request_end` 里的轨迹摘要 + +除了逐个 `chain_step`,`request_end` 还带三个便于做报表的字段: + +| 字段 | 含义 | +|---|---| +| `chain_walk` | 整个遍历的步骤数组(上限 12 步,超出截断) | +| `degraded` | 布尔。`true` = 有过跳过/失败,即**发生了降级** | +| `tier_served` | 实际服务的那一档;直连或全失败时为 `-1` | + +> **计费口径**:按**实际服务的模型**计费。降级到 tier 3 仍按 tier 3 的价算, +> `chain_step` / `degraded` / `tier_served` 只作**观测**,不参与计价。 +> 理由见 §7.5。 + +`request_end` 放在 `writeRec` 是因为四条入口路径(直连/AUTO × 流式/非流式)都 +经过它,既不会漏(流式的 token 数只有流结束才知道),也不会重复。 + +**hot path 注意事项**:没有插件注册某 stage 时,`Fire` 立刻返回(一次 `RLock` +加一次 map 查找)。装了插件之后,每个请求会在该 stage 上多一次 Lua 调用—— +这是同步的,在关键路径上。计费插件那种"每请求一次"是正常的;把重活放进钩子是 +反模式。 + +### 3.2 钩子的返回值 + +- 返回 `nil` 或不返回 = **没有意见**,payload 原样传给下一个插件 +- 返回 table = 其中的键会**合并进 payload**,并作为 `Fire` 的返回值 + +**同一 stage 的多个插件并行执行**,但返回值按**插件加载顺序**合并,所以结果是 +确定的(不依赖 goroutine 调度)。代价是一个插件看不到另一个插件刚加的字段: +每个钩子拿到的是**同一份 payload 快照**。 + +这与早期版本不同 —— 早期是顺序执行,后一个插件能看到前一个的返回值。它从未被 +实际依赖(随核心发布的 billing 在每个 stage 都 `return nil`,注释里写着 +"nobody downstream would read a return value"),但这是一处**契约变化**:如果你的 +插件依赖「读到前一个插件写的字段」,并行的两个插件之间必须改用外部通信 +(例如各自写 `plugin.state`,由 `/api/plugins//state` 读取)。 + +单个插件时不启 goroutine,直接调用。 + +前三个 stage 的返回值目前没有内部消费者(最后一个 stage 之后就是写审计), +所以计费插件改用 `plugin.state` + `/state` 端点来暴露数据。 + +--- + +## 4. 界面扩展 + +### 4.1 整页 + +```lua +plugin.ui = { + page = { + page_id = "billing", -- 必填,kebab-case。侧栏 data-tab 与 #tab-billing + title = "Billing", -- 必填,侧栏文字 + icon = "💰", -- 可选 + order = 40, -- 侧栏排序,默认 100 + mount = [[...HTML...]], + }, +} +``` + +### 4.2 往现有页面追加元素 + +```lua +plugin.ui = { + elements = { + { + target = "status", -- status | chat | keys | sort | sources | adapters + anchor = "top", -- "top" | "bottom" | "before:" | "after:" + order = 5, + mount = [[...HTML...]], + }, + }, +} +``` + +### 4.3 `mount` 里可以带 ` diff --git a/internal/gateway/ui_contract_test.go b/internal/gateway/ui_contract_test.go index cfa5d50..332bbb0 100644 --- a/internal/gateway/ui_contract_test.go +++ b/internal/gateway/ui_contract_test.go @@ -59,6 +59,82 @@ func lineOf(src string, idx int) int { return strings.Count(src[:idx], "\n") + 1 } +// TestUIEditFormRoundTripsEveryOptionalSourceField pins the JS half of the +// source-edit contract. +// +// The server distinguishes "absent" from "explicitly empty" for proxy_url, +// api_key_env and timeout (they are pointers). That only helps if the form +// SENTS them — a form that omits them falls back to "inherit", so clearing the +// proxy box would silently keep the old proxy, which is the exact bug in the +// other direction. +// +// It must also send the api_key_env value it was given, since that is the one +// field whose loss is invisible until the next upstream call. +func TestUIEditFormRoundTripsEveryOptionalSourceField(t *testing.T) { + src := uiSource(t) + body, ok := jsFunctionBody(src, "saveSource") + if !ok { + t.Fatal("saveSource() not found in the WebUI") + } + for _, field := range []string{"proxy_url", "api_key_env", "timeout"} { + if !strings.Contains(body, field+":") { + t.Errorf("saveSource() does not send %q; the server treats an absent "+ + "field as \"keep the stored value\", so the form could never clear it", field) + } + } + // The form's inputs must be filled from the values the server sends, or a + // save would post an empty box and clear a configured field. + for _, read := range []string{`$("#s-proxy")`, `$("#s-keyenv")`, `$("#s-timeout")`} { + if !strings.Contains(src, read+".value") { + t.Errorf("the source form never reads %s — it would post an empty value", read) + } + } +} + +// TestUIEditFormReadsDurationsFromTheRevealView: config.Source tags +// Timeout/QueueTimeout `json:"-"`, so the durations only reach the form if the +// reveal endpoint adds them explicitly. The form MUST read them (see above), +// which makes this API response load-bearing rather than informational. +func TestUIEditFormReadsDurationsFromTheRevealView(t *testing.T) { + src := uiSource(t) + for _, id := range []string{"s-proxy", "s-keyenv", "s-timeout"} { + if !strings.Contains(src, `id="`+id+`"`) { + t.Errorf("source form input #%s is missing", id) + } + } + // The timeout box must render through the duration helper. Reading a raw + // time.Duration number would show nanoseconds in a box that expects "300s". + if !strings.Contains(src, "timeoutText(") { + t.Error("the timeout input does not go through timeoutText(); a raw " + + "time.Duration JSON number is nanoseconds and would be uneditable") + } +} + +// TestUIEditFormSendsEverySourceField is the mirror of +// TestSourcePayloadCoversEveryEditableField on the JavaScript side: each field +// the server accepts must be submitted by the form, or the server's upsert has +// nothing to store for it. +func TestUIEditFormSendsEverySourceField(t *testing.T) { + src := uiSource(t) + body, ok := jsFunctionBody(src, "saveSource") + if !ok { + t.Fatal("saveSource() not found") + } + // Fields the edit dialog owns. proxy_url / api_key_env / timeout are + // covered by the round-trip test above; this one catches the rest. + // models and meta are local variables submitted by Go shorthand + // (payload = { models, meta }), so they are matched as bare identifiers. + for _, field := range []string{ + "name:", "base_url:", "api_key:", "adapter:", "endpoint:", + "image_endpoint:", "max_concurrent:", "rpm:", "temperature:", + "models,", "meta,", + } { + if !strings.Contains(body, field) { + t.Errorf("saveSource() does not send %s", field) + } + } +} + // TestUIAPICallsDeclareMethod asserts that every api() call passing an options // object also declares an HTTP method (or is a GET that only passes an // AbortSignal). Without this, fetch defaults to GET and mutating endpoints are diff --git a/internal/gateway/ui_plugin_test.go b/internal/gateway/ui_plugin_test.go new file mode 100644 index 0000000..0a5e2a2 --- /dev/null +++ b/internal/gateway/ui_plugin_test.go @@ -0,0 +1,390 @@ +package gateway + +import ( + "strings" + "testing" +) + +// The plugin UI injection is JavaScript inside the embedded index.html, and it +// is the ONLY thing that turns a plugin's `ui` block into a visible page or +// element. These tests pin the wiring on the JS side; the server side (what the +// payload contains) is covered by TestUIInjectServesPluginUI and the lua +// package's TestBillingPluginDeclaresUI. +// +// What makes this worth pinning: a missing hook here fails SILENTLY. The page +// simply never appears, there is no error anywhere, and it looks like "the +// plugin didn't declare a page" rather than "the UI forgot to inject it". + +// uiSource returns the embedded WebUI document. +func uiSourceX(t *testing.T) string { + t.Helper() + return uiSource(t) +} + +// TestUIFetchesPluginInjection: the boot sequence must ask the kernel what to +// inject. Without this fetch the whole feature is inert. +func TestUIFetchesPluginInjection(t *testing.T) { + src := uiSourceX(t) + if !strings.Contains(src, "/api/ui-inject") { + t.Error("the WebUI never calls /api/ui-inject; plugin pages and elements can never appear") + } +} + +// TestUIInjectsBeforeFirstRender: injection must be awaited before the first +// refresh, otherwise the sidebar is built without the plugin entry and the +// first paint races the fetch. This is an ordering contract, so it is asserted +// on the source order rather than trusted. +func TestUIInjectsBeforeFirstRender(t *testing.T) { + src := uiSourceX(t) + iInject := strings.Index(src, "injectPluginUI()") + iRefresh := strings.LastIndex(src, `refresh("status")`) + if iInject < 0 { + t.Fatal("injectPluginUI() is never called") + } + if iRefresh < 0 { + t.Fatal("the boot sequence no longer calls refresh(\"status\")") + } + if iInject > iRefresh { + t.Error("injectPluginUI() is called after the first refresh; the sidebar " + + "and #main would be built before the plugin page exists") + } + // And it must be awaited, not fire-and-forget. + window := src[iInject:] + if !strings.Contains(window[:200], ".finally") && !strings.Contains(window[:200], "await") { + t.Error("injectPluginUI() is not awaited before refresh; a slow response " + + "would race the first paint") + } +} + +// TestUIPluginScriptsRunAfterMarkup is the subtle one. Setting innerHTML with a +// `) + var out []string + for _, m := range re.FindAllStringSubmatch(html, -1) { + out = append(out, m[1]) + } + return out +} + +// TestBillingPageCoversItsData is the other half: the page declares tables for +// per-source / per-model / per-key / per-day and must actually render into them. +// A table id that is never written to is exactly how "the page loads and shows +// nothing" happens without an error. +func TestBillingPageCoversItsData(t *testing.T) { + page, _ := billingUI(t) + for _, id := range []string{ + "billing-kpis", "billing-by-source", "billing-by-model", + "billing-by-key", "billing-by-day", + } { + if !strings.Contains(page, `id="`+id+`"`) { + t.Errorf("the page has no container #%s", id) + } + } + // Every container must be written to by the script, not just declared. + scripts := extractScripts(page) + all := strings.Join(scripts, "\n") + for _, id := range []string{"billing-kpis", "billing-by-source", "billing-by-model", "billing-by-key", "billing-by-day"} { + if !strings.Contains(all, `getElementById("`+id+`")`) { + t.Errorf("#%s is declared but never read by the script — it stays empty forever", id) + } + } +} + +// TestPluginIconIsNotARawEmoji guards the sidebar icon. The billing plugin +// declared icon = "💰" and the WebUI drops that verbatim into the button, while +// every native tab uses an inline SVG styled with `stroke: currentColor`. An +// emoji there renders at the wrong size and ignores the theme, so it does not +// match its neighbours — which is what the operator reported. +func TestPluginIconIsNotARawEmoji(t *testing.T) { + icon := billingIcon(mustBillingSource(t)) + if icon == "" { + t.Fatal("billing declares no icon; the sidebar entry would be blank") + } + if isEmojiIcon(icon) { + t.Errorf("page icon is the raw emoji %q — the WebUI sidebar renders "+ + "native tabs as inline SVG (stroke: currentColor), so an emoji is the "+ + "wrong size and ignores the theme. Use an inline SVG path instead.", icon) + } +} + +// isEmojiIcon reports whether s is a pictographic emoji rather than markup or a +// text glyph. Codepoints in the pictographic blocks, plus the regional-indicator +// pair used by flags. +func isEmojiIcon(s string) bool { + r := []rune(s) + if len(r) == 0 { + return false + } + // Anything containing '<' is markup (an inline ), which is the fix. + if strings.ContainsRune(s, '<') { + return false + } + for _, c := range r { + switch { + case c >= 0x1F300 && c <= 0x1FAFF, // pictographs, symbols, supplemental + c >= 0x1F000 && c <= 0x1F2FF, // mahjong/domino/cards + c >= 0x2600 && c <= 0x27BF, // misc symbols + dingbats + c >= 0x2B00 && c <= 0x2BFF, // arrows/misc symbols + c == 0xFE0F, // variation selector-16 + c >= 0x1F1E6 && c <= 0x1F1FF: // regional indicators (flags) + return true + } + } + return false +} + +func mustBillingSource(t *testing.T) string { + t.Helper() + b, err := os.ReadFile("plugins/billing.lua") + if err != nil { + t.Fatal(err) + } + return string(b) +} + +// billingIcon extracts the declared page icon. It accepts BOTH a quoted string +// and a [==[ ... ]==] long string, because an inline SVG cannot be written as a +// Lua short string without escaping every quote in it. +func billingIcon(src string) string { + if m := regexp.MustCompile(`(?m)^\s*icon\s*=\s*"([^"]*)"`).FindStringSubmatch(src); m != nil { + return m[1] + } + if m := regexp.MustCompile(`(?s)\bicon\s*=\s*\[==\[(.*?)\]==\]`).FindStringSubmatch(src); m != nil { + return m[1] + } + return "" +} + +// TestBillingTableHeaderMatchesRowColumns catches column drift. +// +// row() gained cache columns (fresh / cache / cache%) while the header row did +// not, in the same edit. The result is a table whose cells are shifted one +// column left from "fresh" onward — so "cache%" sits under "completion" and the +// last cell has no label. It renders, it has data, and it is wrong in a way that +// takes a careful read to notice. +func TestBillingTableHeaderMatchesRowColumns(t *testing.T) { + page, _ := billingUI(t) + js := strings.Join(extractScripts(page), "\n") + if !strings.Contains(js, "") { + t.Fatal("no table header found in the billing page script") + } + hStart := strings.Index(js, "function tableFor") + if hStart < 0 { + t.Fatal("no tableFor in the billing page script") + } + // Count the usage table's headers ONLY. The count is taken from tableFor to + // the end of the script, which used to be fine when the rule editor lived on + // its own page; now that both halves share one script, the rule table's nine + // headers were counted too and the check reported 17 vs 8 — a failure about + // two unrelated tables. Stop at the rule editor's script. + header := js[hStart:] + if cut := strings.Index(header, "/api/plugins/billing/rules"); cut > 0 { + header = header[:cut] + } + // Count with OR without attributes: the name column now carries + // style='width:22%' so the header and its data cell can be given the same + // width, and a bare "" count silently reported 7 vs 8 — the check + // failing on the very column it had just been taught to size. + th := strings.Count(header, "") + strings.Count(header, "; the first cell uses .. so + // counting is exact. + td := strings.Count(row, "") + + if th != td { + t.Errorf("★ header declares %d columns but row() emits %d cells — the "+ + "table is misaligned from the first differing column on", th, td) + } +} + +// TestRuleTableLocksColumnWidthsInAColgroup guards the layout class of bug that +// only a screenshot caught: two adjacent headers rendered on top of each other, +// the URL input squeezed to "https:", and the per-model × button landing 78px +// past its cell onto the neighbouring Delete button. +// +// The measurement that found each one: +// - table-layout:fixed with widths on → 954px divided into 9 equal +// 106px columns, every declared width discarded. +// - widths on WITHOUT fixed → browser sizes from content and ignores +// them; the 70px currency column still collapsed to 48px (24px of input). +// - fixed + widths → they take effect, but the model-price cell's width +// then depended on WHICH RULE was widest (352px vs 206px). +// +// The fix is all three together: fixed layout, a colgroup, and flex children +// that may shrink. +func TestRuleTableLocksColumnWidthsInAColgroup(t *testing.T) { + page, _ := billingUI(t) + js := strings.Join(extractScripts(page), "\n") + // Match the ATTRIBUTE, not the bare string: a substring check on + // "table-layout:fixed" is satisfied by the explanatory comment that sits + // three lines above the tag, so removing the attribute from the table still + // passed (mutation-verified). Require it inside a style='...' literal. + if !strings.Contains(js, `font-size:12px;table-layout:fixed'`) { + t.Error("the rule table must stay table-layout:fixed — with auto layout " + + "the widths below are ignored and columns size from content") + } + colStart := strings.Index(js, "") + if colStart < 0 { + t.Fatal("the rule table declares no ; per-column widths on " + + " or are advisory and get recomputed from content") + } + // The colgroup is built from a JS array of pixel widths, so assert on the + // array, not on literal ""`) { + t.Error("colgroup entries are no longer emitted as ") + } + // Every price-row INPUT must be allowed to shrink, or the row's fixed widths + // sum past the cell and the × button lands on the next column. The × button + // is the opposite: it must NOT shrink (flex:0 0 auto), because a squashed + // delete button is unclickable. + for _, needle := range []string{"m-name", "m-p", "m-c"} { + i := strings.Index(js, "class='"+needle+"'") + if i < 0 { + i = strings.Index(js, "class='ghost small "+needle+"'") + } + if i < 0 { + t.Errorf("price row lost the %s input/button", needle) + continue + } + window := js[i : i+220] + if !strings.Contains(window, "flex:") { + t.Errorf("%s has no flex sizing; the row is 80+96+96+38 ≈ 330px wide "+ + "against a 262px cell and overflows onto the next column", needle) + } + if !strings.Contains(window, "min-width:0") { + t.Errorf("%s lacks min-width:0; a flex item will not shrink below its "+ + "content width, which is how the × button ended up 78px past its cell", needle) + } + } + i := strings.Index(js, "class='ghost small m-del'") + if i < 0 { + t.Fatal("the per-model delete button is gone") + } + if !strings.Contains(js[i:i+120], "flex:0 0 auto") { + t.Error("the per-model × button must be flex:0 0 auto — it is a fixed-size " + + "control and must never be squeezed by the flexible inputs") + } +} diff --git a/internal/lua/fastvalue.go b/internal/lua/fastvalue.go new file mode 100644 index 0000000..0388e79 --- /dev/null +++ b/internal/lua/fastvalue.go @@ -0,0 +1,139 @@ +package lua + +import ( + "encoding/json" + "strconv" +) + +// This file exists to remove the JSON round-trip from the plugin hook path. +// +// Measured before this change: one Fire() call cost 14.6us, of which +// json.Marshal(payload) was 4.0us and json.Unmarshal back into interface{} was +// 5.6us — 67% of the whole hook, spent re-deriving a tree that the very next +// line (pushGoValue) already knows how to walk natively. +// +// The round-trip existed because payload arrives as map[string]interface{} and +// pushGoValue wanted a plain interface{} tree. pushGoValue already handles +// map[string]interface{} and []interface{} directly, so the fix is a converter +// that flattens the concrete Go types the gateway actually produces, instead of +// a generic serialize/parse. +// +// The JSON path is kept as the fallback for types the fast path does not know, +// so a plugin payload carrying something exotic still arrives rather than being +// silently dropped. + +// jsonNumberType is the type json.Unmarshal produces for every JSON number. +// The fast path recognizes it so numbers survive as numbers rather than being +// stringified into a different Lua type. +type jsonNumberType = float64 + +// fastToPlain converts the concrete Go values the gateway puts in a hook +// payload into the interface{} tree pushGoValue consumes, without JSON. +// +// The interesting cases are the ones json.Unmarshal would have normalized: +// - json.Number-ish types (all float64 in practice) +// - int / int64 / uint variants, which must land as Lua numbers, not as the +// strings a naive "only handle float64" switch would produce +// - []string and map[string]string, extremely common in payloads and NOT +// handled by a switch that only knows []interface{} +// +// Anything unrecognized returns (value, false) so the caller can fall back to +// the JSON path, which is slower but total. +func fastToPlain(v interface{}) (interface{}, bool) { + switch x := v.(type) { + case nil, bool, string, jsonNumberType: + return v, true + case int: + return float64(x), true + case int8: + return float64(x), true + case int16: + return float64(x), true + case int32: + return float64(x), true + case int64: + return float64(x), true + case uint: + return float64(x), true + case uint8: + return float64(x), true + case uint16: + return float64(x), true + case uint32: + return float64(x), true + case uint64: + return float64(x), true + case float32: + return float64(x), true + case []string: + out := make([]interface{}, len(x)) + for i, s := range x { + out[i] = s + } + return out, true + case map[string]string: + out := make(map[string]interface{}, len(x)) + for k, s := range x { + out[k] = s + } + return out, true + case map[string]interface{}: + // The overwhelmingly common case, and the one the hook path always + // takes at the top level. + out := make(map[string]interface{}, len(x)) + for k, it := range x { + c, ok := fastToPlain(it) + if !ok { + return nil, false + } + out[k] = c + } + return out, true + case []interface{}: + out := make([]interface{}, len(x)) + for i, it := range x { + c, ok := fastToPlain(it) + if !ok { + return nil, false + } + out[i] = c + } + return out, true + default: + return nil, false + } +} + +// plainForLua returns the interface{} tree to hand to pushGoValue, using the +// fast path when it can and JSON only when it must. +func plainForLua(payload map[string]interface{}) interface{} { + if fast, ok := fastToPlain(payload); ok { + return fast + } + // Rare: some type the fast path does not model. Serialize and let the + // generic decoder normalize it, so the hook still sees the field. + b, err := json.Marshal(payload) + if err != nil { + return payload + } + var decoded interface{} + if err := json.Unmarshal(b, &decoded); err != nil { + return payload + } + return decoded +} + +// toPlainSlice is the []interface{} entry point of fastToPlain, exposed so the +// Fire path can flatten a pre-built slice without re-walking the map header. +func toPlainSlice(v []interface{}) []interface{} { + if fast, ok := fastToPlain(v); ok { + return fast.([]interface{}) + } + return v +} + +// luaNumberString renders a float the way Lua would print it, for the rare case +// a hook wants the textual form. Kept out of the hot path. +func luaNumberString(f float64) string { + return strconv.FormatFloat(f, 'g', -1, 64) +} diff --git a/internal/lua/fastvalue_test.go b/internal/lua/fastvalue_test.go new file mode 100644 index 0000000..75d47f0 --- /dev/null +++ b/internal/lua/fastvalue_test.go @@ -0,0 +1,132 @@ +package lua + +import ( + "encoding/json" + "reflect" + "testing" +) + +// The hook payload used to be JSON round-tripped on every call. It no longer is, +// so the fast path and the old JSON path must be indistinguishable — otherwise +// a plugin silently sees a different payload than before, which is the worst +// kind of change: it compiles, passes a smoke test, and misprices traffic. +// +// These tests therefore compare the two paths on the SAME inputs rather than +// asserting the fast path's output in isolation. + +func viaJSON(t *testing.T, payload map[string]interface{}) interface{} { + t.Helper() + b, err := json.Marshal(payload) + if err != nil { + t.Fatalf("marshal: %v", err) + } + var decoded interface{} + if err := json.Unmarshal(b, &decoded); err != nil { + t.Fatalf("unmarshal: %v", err) + } + return decoded +} + +// payloads that exercise every branch of the fast converter. +func payloadCases() []map[string]interface{} { + return []map[string]interface{}{ + {"model": "deepseek-v4.1-flash", "source": "commandcode", "ok": true}, + // Numbers of every width: a converter that only knows float64 turns + // ints into something else, and a token count that arrives as a string + // makes a Lua hook do arithmetic on nil. + {"i": 42, "i8": int8(8), "i16": int16(16), "i32": int32(32), "i64": int64(1 << 40), + "u": uint(7), "u64": uint64(1 << 50), "f32": float32(1.5), "f64": 2.25}, + // Zero, negative, and very large values must stay numbers. + {"zero": 0, "neg": -17, "huge": 1e308, "tiny": 1e-308}, + // Slices and string maps: common in payloads and absent from a switch + // that only knows []interface{} / map[string]interface{}. + {"msgs": []interface{}{"a", "b"}, "tags": []string{"x", "y"}}, + {"kv": map[string]string{"a": "1", "b": "2"}}, + // Nesting, which is where a shallow converter silently drops a level. + {"usage": map[string]interface{}{ + "prompt_tokens": 2048, "cache_hit_tokens": 1024, + "nested": map[string]interface{}{"deep": []interface{}{1, "two", true, nil}}, + }}, + {"nil_field": nil, "empty_map": map[string]interface{}{}, "empty_slice": []interface{}{}}, + // A value the fast path does NOT model: it must fall back to JSON and + // still arrive, not disappear. + {"weird": struct { + A int `json:"a"` + B string `json:"b"` + }{1, "x"}}, + } +} + +func TestFastPathMatchesJSONPath(t *testing.T) { + for i, p := range payloadCases() { + want := viaJSON(t, p) + got := plainForLua(p) + if !reflect.DeepEqual(want, got) { + t.Errorf("case %d: fast path differs from JSON path\n payload: %#v\n json: %#v\n fast: %#v", + i, p, want, got) + } + } +} + +// TestFastPathIsActuallyUsed guards against the fast path silently degrading to +// JSON for the payload the gateway really sends. If a future payload gains a +// type the converter does not model, this still works (it falls back) but the +// optimization is gone — and the next person measuring the hook would be +// measuring the old cost without knowing why. +func TestFastPathIsActuallyUsed(t *testing.T) { + // This mirrors the real request_end payload shape from the gateway. + realistic := map[string]interface{}{ + "stage": "request_end", "kind": "end", "model": "deepseek-v4.1-flash", + "source": "commandcode", "key": "stress-key", "ok": true, + "status": 200, "duration_ms": 1234, + "usage": map[string]interface{}{ + "prompt_tokens": float64(2048), "completion_tokens": float64(512), + "cache_hit_tokens": float64(1024), "total_tokens": float64(3584), + }, + "walk": []interface{}{ + map[string]interface{}{"kind": "tier_skip", "tier": 1, "source": "", "model": "", "reason": "cooldown"}, + map[string]interface{}{"kind": "selected", "tier": 2, "source": "commandcode", "model": "m", "reason": ""}, + }, + "ts": float64(1700000000), + } + if _, ok := fastToPlain(realistic); !ok { + t.Errorf("★ the realistic request payload does NOT take the fast path — " + + "the JSON round-trip is still on the hot path for real traffic") + } +} + +// TestFastPathDoesNotAliasInput: the converter builds a new tree. If it ever +// returned the caller's map directly, a Lua hook's writes could not reach Go — +// but worse, a later mutation of the payload would race with the snapshot the +// persistence saver is holding. +func TestFastPathDoesNotAliasInput(t *testing.T) { + src := map[string]interface{}{ + "usage": map[string]interface{}{"prompt_tokens": float64(1)}, + "list": []interface{}{"a"}, + } + out, ok := fastToPlain(src) + if !ok { + t.Fatal("fast path declined a plain payload") + } + m := out.(map[string]interface{}) + m["new"] = "added" + src["also_new"] = "must not appear" + + if _, leaked := m["also_new"]; leaked { + t.Error("output map aliases the input map") + } + if _, leaked := src["new"]; leaked { + t.Error("writing to the output mutated the input") + } + // And the nested maps must be copies too. + inner := m["usage"].(map[string]interface{}) + inner["prompt_tokens"] = float64(999) + if src["usage"].(map[string]interface{})["prompt_tokens"] != float64(1) { + t.Error("nested map is shared, not copied — a hook could mutate the payload") + } + list := m["list"].([]interface{}) + list[0] = "changed" + if src["list"].([]interface{})[0] != "a" { + t.Error("nested slice is shared, not copied") + } +} diff --git a/internal/lua/fire_parallel_test.go b/internal/lua/fire_parallel_test.go new file mode 100644 index 0000000..7480a1e --- /dev/null +++ b/internal/lua/fire_parallel_test.go @@ -0,0 +1,334 @@ +package lua + +import ( + "os" + "path/filepath" + "sync" + "testing" + "time" +) + +// Fire runs same-stage hooks in parallel across plugins. These tests pin the +// three things that make that safe, each of which failed at least once while +// the change was being written. + +// twoHookPlugins loads two plugins that both hook request_end. +func twoHookPlugins(t *testing.T, codeA, codeB string) *Plugins { + t.Helper() + vm := NewVM(filepath.Join(t.TempDir(), "adapters")) + if err := vm.Start(); err != nil { + t.Fatalf("vm start: %v", err) + } + t.Cleanup(vm.Stop) + ps := NewPlugins(vm, filepath.Join(t.TempDir(), "plugins")) + ps.save.setWakeDelayForTest(time.Millisecond) + t.Cleanup(ps.Close) + ps.DisableStatePersistence() + if err := ps.LoadSource("aaa", codeA); err != nil { + t.Fatalf("load aaa: %v", err) + } + if err := ps.LoadSource("bbb", codeB); err != nil { + t.Fatalf("load bbb: %v", err) + } + return ps +} + +func counterPlugin(name string, key string) string { + // Built by concatenation rather than fmt: %q is not a Go verb, and the first + // attempt produced Lua source with a stray '=' that only failed at compile + // time inside three different tests. + return ` +local p = { name = "` + name + `" } +p.state = { n = 0 } +p.hooks = { request_end = "bump" } +function p.bump(payload) + p.state.n = p.state.n + 1 + if "` + key + `" ~= "" then return { who = "from-` + name + `" } end + return nil +end +return p +` +} + +// TestSingleAndParallelPathsMergeIdentically: the single-plugin fast path and +// the multi-plugin parallel path must produce the same payload. +// +// This is not theoretical. The parallel path was written first and the +// single-plugin path kept its old shape; the merge loop was left off the fast +// path, so a lone plugin returning a table had its return value DISCARDED. The +// existing stage-order test caught it — but only because it happened to check +// the payload after Fire. A plugin that returned fields nobody read would have +// broken silently. +func TestSingleAndParallelPathsMergeIdentically(t *testing.T) { + single := twoHookPlugins(t, counterPlugin("aaa", "who"), ` +local p = { name = "zzz" } +p.state = { n = 0 } +p.hooks = { request_end = "bump" } +function p.bump(payload) p.state.n = p.state.n + 1 return nil end +return p +`) + // remove the second so exactly one plugin hooks this stage + single.mu.Lock() + single.plugins = single.plugins[:1] + single.mu.Unlock() + single.rebuild() + + ps := twoHookPlugins(t, counterPlugin("aaa", "who"), counterPlugin("bbb", "who")) + + singleOut := single.Fire(StageRequestEnd, map[string]interface{}{"model": "m"}) + psOut := ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m"}) + + if singleOut["who"] != "from-aaa" { + t.Errorf("★ single-plugin path dropped the return value: %v", singleOut) + } + // Two plugins both return `who`; the merge order decides the winner, and it + // must be deterministic (load order), not scheduler-dependent. + if psOut["who"] != "from-bbb" { + t.Errorf("merge order is not load order: who = %v, want from-bbb", psOut["who"]) + } +} + +// TestMergeOrderIsDeterministic: goroutine completion order must not leak into +// the result. Running Fire repeatedly must always yield the same payload, or a +// gateway's behaviour changes run to run with the same plugins installed. +func TestMergeOrderIsDeterministic(t *testing.T) { + // The plugins differ in cost so a scheduler-dependent merge is visible. + slow := ` +local p = { name = "aaa" } +p.state = { n = 0 } +p.hooks = { request_end = "slow" } +function p.slow(payload) + local acc = 0 + for i = 1, 3000 do acc = acc + i % 7 end + p.state.n = p.state.n + 1 + return { winner = "aaa", acc = acc } +end +return p +` + fast := ` +local p = { name = "bbb" } +p.state = { n = 0 } +p.hooks = { request_end = "fast" } +function p.fast(payload) + p.state.n = p.state.n + 1 + return { winner = "bbb" } +end +return p +` + ps := twoHookPlugins(t, slow, fast) + first := ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m"})["winner"] + for i := 0; i < 60; i++ { + got := ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m"})["winner"] + if got != first { + t.Fatalf("★ winner changed between runs: %q then %q — merge depends on goroutine "+ + "scheduling, so the same configuration behaves differently run to run", first, got) + } + } + // Load order decides: bbb is loaded second, so it wins. + if first != "bbb" { + t.Errorf("winner = %q, want bbb (load order decides, not timing)", first) + } +} + +// TestThrowingHookDoesNotBreakSiblings covers the failure path that golua +// actually produces: a hook raising a Lua error. +// +// The test name used to claim it covered panics, and it used error() to do it. +// Probing the six ways a Lua program can fault (error(), indexing nil, calling +// nil, concatenating nil, arithmetic on nil, unbounded recursion) showed golua +// converts ALL of them into an error RETURN, not a Go panic — so the recover() +// in Fire was untested by that case, and deleting recover() still passed. The +// test was renamed to say what it verifies. +// +// recover() is kept anyway: it guards the Go side of Fire (a nil map write, a +// future change to how the payload is prepared), which is cheap and cannot be +// triggered from Lua today. Claiming it is covered by a Lua test would be the +// kind of assurance that evaporates the first time someone checks. +func TestThrowingHookDoesNotBreakSiblings(t *testing.T) { + throwing := ` +local p = { name = "aaa" } +p.state = {} +p.hooks = { request_end = "boom" } +function p.boom(payload) error("intentional failure") end +return p +` + healthy := counterPlugin("bbb", "") + ps := twoHookPlugins(t, throwing, healthy) + + out := ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m"}) // must not crash + if out == nil { + t.Error("Fire returned nil") + } + st, ok := ps.State("bbb").(map[string]interface{}) + if !ok { + t.Fatal("healthy plugin has no state") + } + if n := st["n"]; n != float64(1) { + t.Errorf("★ healthy plugin did not run alongside the throwing one: n = %v", n) + } + he := ps.HookErrors() + if len(he) == 0 { + t.Error("a throwing hook was not recorded in hook_errors — the failure would be invisible") + } +} + +// TestParallelHooksAllRunOnce: every plugin must be invoked exactly once per +// Fire. A lost or duplicated goroutine shows up as a wrong total, which for the +// billing plugin means a wrong bill. +func TestParallelHooksAllRunOnce(t *testing.T) { + ps := twoHookPlugins(t, counterPlugin("aaa", ""), counterPlugin("bbb", "")) + const fires = 100 + for i := 0; i < fires; i++ { + ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m"}) + } + for _, name := range []string{"aaa", "bbb"} { + st, ok := ps.State(name).(map[string]interface{}) + if !ok { + t.Fatalf("%s has no state", name) + } + if n := st["n"]; n != float64(fires) { + t.Errorf("%s counted %v hooks, want %d", name, n, fires) + } + } +} + +// TestConcurrentFireIsSafe drives Fire from many goroutines at once. Each +// plugin has its own Lua state and its own mutex, so this must hold; the race +// detector is what proves it, not the assertions. +func TestConcurrentFireIsSafe(t *testing.T) { + ps := twoHookPlugins(t, counterPlugin("aaa", ""), counterPlugin("bbb", "")) + var wg sync.WaitGroup + for g := 0; g < 8; g++ { + wg.Add(1) + go func() { + defer wg.Done() + for i := 0; i < 50; i++ { + ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "source": "s"}) + } + }() + } + wg.Wait() + for _, name := range []string{"aaa", "bbb"} { + st, ok := ps.State(name).(map[string]interface{}) + if !ok { + t.Fatalf("%s has no state", name) + } + if n := st["n"]; n != float64(400) { + t.Errorf("%s counted %v, want 400", name, n) + } + } +} + +// TestFireDoesNotHoldRegistryLockWhileRunningHooks: a hook that installs or +// disables a plugin takes ps.mu for write. If Fire held ps.mu across the hook, +// that would self-deadlock — the exact failure TrueAgent recorded for its own +// parallel stop path. +func TestFireDoesNotHoldRegistryLockWhileRunningHooks(t *testing.T) { + reentrant := ` +local p = { name = "aaa" } +p.state = {} +p.hooks = { request_end = "reenter" } +function p.reenter(payload) + -- Reading the registry from inside a hook is the read half of the same lock. + local _ = #payload + return nil +end +return p +` + vm := NewVM(filepath.Join(t.TempDir(), "adapters")) + if err := vm.Start(); err != nil { + t.Fatal(err) + } + defer vm.Stop() + ps := NewPlugins(vm, filepath.Join(t.TempDir(), "plugins")) + defer ps.Close() + ps.DisableStatePersistence() + if err := ps.LoadSource("aaa", reentrant); err != nil { + t.Fatal(err) + } + if err := ps.LoadSource("bbb", counterPlugin("bbb", "")); err != nil { + t.Fatal(err) + } + + done := make(chan struct{}) + go func() { + defer close(done) + ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m"}) + // Touch the registry the way an admin request would, right after. + _ = ps.Count() + _ = ps.List() + }() + select { + case <-done: + case <-time.After(10 * time.Second): + t.Fatal("Fire deadlocked against the registry lock") + } +} + +// TestBillingSurvivesParallelFire runs the REAL bundled plugin next to another +// one. A synthetic counter cannot catch a mismatch between the documented +// payload shape and what the plugin actually reads. +func TestBillingSurvivesParallelFire(t *testing.T) { + dir := t.TempDir() + pdir := filepath.Join(dir, "plugins") + os.MkdirAll(pdir, 0o755) + src, err := os.ReadFile("plugins/billing.lua") + if err != nil { + t.Fatalf("read billing.lua: %v", err) + } + if err := os.WriteFile(filepath.Join(pdir, "billing.lua"), src, 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(pdir, "observer.lua"), + []byte(`local p={name="observer"} p.state={n=0} +p.hooks={request_end="see"} +function p.see(payload) p.state.n=p.state.n+1 return nil end +return p`), 0o644); err != nil { + t.Fatal(err) + } + + vm := NewVM(filepath.Join(dir, "adapters")) + if err := vm.Start(); err != nil { + t.Fatal(err) + } + defer vm.Stop() + ps := NewPlugins(vm, pdir) + ps.save.setWakeDelayForTest(time.Millisecond) + defer ps.Close() + ps.DisableStatePersistence() + if err := ps.LoadDir(); err != nil { + t.Fatal(err) + } + if ps.Count() != 2 { + t.Fatalf("loaded %d plugins, want 2", ps.Count()) + } + + payload := map[string]interface{}{ + "model": "deepseek-v4.1-flash", "source": "commandcode", "key": "k", "ok": true, + "usage": map[string]interface{}{ + "prompt_tokens": float64(1000), "completion_tokens": float64(100), + "cache_hit_tokens": float64(0), + }, + } + ps.Fire(StageRequestEnd, payload) + + st := ps.State("billing") + if st == nil { + t.Fatal("billing produced no state") + } + m := st.(map[string]interface{}) + total, ok := m["total"].(map[string]interface{}) + if !ok { + t.Fatalf("billing.total is %T", m["total"]) + } + if total["requests"] != float64(1) { + t.Errorf("billing counted %v requests, want 1 — the plugin was starved by the parallel path", + total["requests"]) + } + if n := ps.State("observer").(map[string]interface{})["n"]; n != float64(1) { + t.Errorf("observer counted %v, want 1", n) + } + if len(ps.HookErrors()) != 0 { + t.Errorf("hook errors under parallel Fire: %v", ps.HookErrors()) + } +} diff --git a/internal/lua/hook_guard_test.go b/internal/lua/hook_guard_test.go new file mode 100644 index 0000000..3bab5b1 --- /dev/null +++ b/internal/lua/hook_guard_test.go @@ -0,0 +1,76 @@ +package lua + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +// A plugin hook that RAISES must not take the process down. This is the +// production outage: a Lua error anywhere in a hook made golua's callEx call +// L.StackTrace(), which calls lua_getinfo and SIGSEGVs on a deep enough stack — +// a C-level signal Go cannot recover from, so one buggy plugin killed the whole +// gateway and took every in-flight request with it. +// +// The fix routes hook calls through a Lua-side pcall guard, so the error comes +// back as an ordinary return value. This test fires hooks that raise on purpose +// and asserts three things: the process survives, the failure is RECORDED, and +// a well-behaved plugin on the same state keeps working afterwards (the error +// must not poison the Lua state). +func TestHookThatRaisesDoesNotCrashTheProcess(t *testing.T) { + ps, pdir := billingVM(t) + boom := ` +local plugin = {} +plugin.name = "boom" +plugin.version = "0.1" +function plugin.request_end(payload) + -- Raise on a table index, the exact shape of the billing bug that caused the + -- outage. Deliberately NOT a syntax error: this must load fine and fail only + -- when invoked. + local x = nil + return x.field +end +return plugin` + if err := os.WriteFile(filepath.Join(pdir, "boom.lua"), []byte(boom), 0644); err != nil { + t.Fatal(err) + } + if err := ps.LoadSource("boom", boom); err != nil { + t.Fatalf("load boom: %v", err) + } + + payload := map[string]interface{}{ + "model": "m", "source": "s", "ok": true, + "prompt_tokens": 100, "completion_tokens": 10, "time": 1750000000000, + } + // Fire many times: a single call could pass by luck, but if the error ever + // escapes into golua's C path the process dies and this test never returns. + for i := 0; i < 50; i++ { + ps.Fire(StageRequestEnd, payload) + } + // Reaching this line at all is the primary assertion. + + errs := ps.HookErrors() + end, ok := errs[string(StageRequestEnd)] + if !ok { + t.Fatal("a raising hook left no record — failures must be observable, not swallowed") + } + if end["count"] == nil || end["count"].(int) == 0 { + t.Error("hook error count is zero despite 50 raising calls") + } + msg, _ := end["last_error"].(string) + if !strings.Contains(msg, "boom") { + t.Errorf("last_error does not name the offending plugin: %q", msg) + } + + // The billing plugin shares the same Plugins registry and must still work: + // one broken plugin may not disable the others. + st, _ := ps.State("billing").(map[string]interface{}) + if st == nil || st["total"] == nil { + t.Fatalf("the healthy plugin stopped working after another plugin raised") + } + tot, _ := st["total"].(map[string]interface{}) + if tot == nil || tot["requests"] == nil || tot["requests"].(float64) == 0 { + t.Errorf("billing recorded no requests after the raising plugin ran: %v", st["total"]) + } +} diff --git a/internal/lua/persist_test.go b/internal/lua/persist_test.go new file mode 100644 index 0000000..e69436c --- /dev/null +++ b/internal/lua/persist_test.go @@ -0,0 +1,283 @@ +package lua + +import ( + "encoding/json" + "os" + "path/filepath" + "testing" + "time" +) + +// These cover plugin state persistence: a plugin's totals live in the Lua VM, +// which dies with the process. Measured on the real gateway before this change: +// 218,241 requests and 26.2M prompt tokens gone after one systemctl restart. + +func persistTestPlugin() string { + return ` +local plugin = { name = "counter", version = "1.0" } +plugin.state = { n = 0, prices_seen = 0 } +plugin.prices = { rate = 1 } +plugin.hooks = { request_end = "bump" } +function plugin.bump(payload) + plugin.state.n = plugin.state.n + 1 + return nil +end +return plugin +` +} + +// newPersistVM wires a plugin registry on a fresh dir, mirroring newPluginVM. +func newPersistVM(t *testing.T) (*Plugins, string) { + t.Helper() + dir := t.TempDir() + vm := NewVM(filepath.Join(dir, "adapters")) + if err := vm.Start(); err != nil { + t.Fatalf("vm start: %v", err) + } + t.Cleanup(vm.Stop) + pdir := filepath.Join(dir, "plugins") + ps := NewPlugins(vm, pdir) + // Shorten the saver debounce: Close() waits for the saver goroutine, so + // with the production 2s interval every test would pay 2s on teardown. + ps.save.setWakeDelayForTest(2 * time.Millisecond) + t.Cleanup(ps.Close) + return ps, pdir +} + +func loadCounter(t *testing.T, ps *Plugins) { + t.Helper() + if err := ps.LoadSource("counter", persistTestPlugin()); err != nil { + t.Fatalf("load: %v", err) + } +} + +func stateN(t *testing.T, ps *Plugins) float64 { + t.Helper() + st := ps.State("counter") + m, ok := st.(map[string]interface{}) + if !ok { + t.Fatalf("state is %T, want map", st) + } + n, ok := m["n"].(float64) + if !ok { + t.Fatalf("state.n is %T, want float64", m["n"]) + } + return n +} + +// TestStateSurvivesRestart is the defect itself: a rebuilt registry over the +// same plugin dir must come up with the previous totals, not at zero. +func TestStateSurvivesRestart(t *testing.T) { + ps, pdir := newPersistVM(t) + loadCounter(t, ps) + for i := 0; i < 25; i++ { + ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "source": "s"}) + } + if got := stateN(t, ps); got != 25 { + t.Fatalf("in-process total = %v, want 25", got) + } + // Force the write the saver would do, so the test does not depend on timing. + ps.save.flush() + ps.Close() + + // A brand-new registry over the same dir: this is the restart. + vm := NewVM(filepath.Join(t.TempDir(), "adapters")) + if err := vm.Start(); err != nil { + t.Fatal(err) + } + defer vm.Stop() + ps2 := NewPlugins(vm, pdir) + defer ps2.Close() + loadCounter(t, ps2) + + if got := stateN(t, ps2); got != 25 { + t.Fatalf("★ total after restart = %v, want 25 — this is the 'restart loses the books' bug", got) + } +} + +// TestPricesAndStateAreSeparate: restoring configuration over history (or the +// reverse) would either erase the totals or resurrect stale prices. +func TestPricesAndStateAreSeparate(t *testing.T) { + ps, pdir := newPersistVM(t) + loadCounter(t, ps) + for i := 0; i < 7; i++ { + ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "source": "s"}) + } + ps.save.flush() + ps.Close() + + b, err := os.ReadFile(ps.stateFile("counter")) + if err != nil { + t.Fatalf("no state file: %v", err) + } + var saved struct { + State map[string]interface{} `json:"state"` + Prices map[string]interface{} `json:"prices"` + } + if err := json.Unmarshal(b, &saved); err != nil { + t.Fatalf("state file is not valid JSON: %v", err) + } + if saved.Prices == nil || saved.Prices["rate"] != float64(1) { + t.Errorf("prices were not persisted separately: %v", saved.Prices) + } + if saved.State["n"] != float64(7) { + t.Errorf("state.n = %v, want 7", saved.State["n"]) + } + _ = pdir +} + +// TestPricesOnlyUpdatePersists guards a real hole: SetState returns early for a +// prices-only payload (correctly leaving state alone), and the first version +// returned before the persistence write — so a reprice was durable in memory +// only and a restart silently reverted to the old prices. +func TestPricesOnlyUpdatePersists(t *testing.T) { + ps, _ := newPersistVM(t) + loadCounter(t, ps) + for i := 0; i < 5; i++ { + ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "source": "s"}) + } + ps.save.flush() + + err := ps.SetState("counter", map[string]interface{}{ + "prices": map[string]interface{}{"rate": 42}, + }) + if err != nil { + t.Fatalf("SetState: %v", err) + } + if got := stateN(t, ps); got != 5 { + t.Errorf("a prices-only payload must not touch state: n = %v, want 5", got) + } + + b, err := os.ReadFile(ps.stateFile("counter")) + if err != nil { + t.Fatalf("no state file after a prices-only update: %v", err) + } + var saved struct { + State map[string]interface{} `json:"state"` + Prices map[string]interface{} `json:"prices"` + } + if err := json.Unmarshal(b, &saved); err != nil { + t.Fatal(err) + } + if saved.Prices["rate"] != float64(42) { + t.Errorf("★ new price not persisted: %v — a restart would revert to the old price", saved.Prices) + } + if saved.State["n"] != float64(5) { + t.Errorf("state was clobbered by a price change: %v", saved.State["n"]) + } +} + +// TestHookReturnValueSurvivesSnapshot guards the ordering bug: snapshotting +// plugin.state resets the Lua stack, so doing it BEFORE reading the hook's +// return value silently turned every opinionated plugin into a silent one — +// breaking the documented "return a table to merge into payload" contract with +// no error anywhere. +func TestHookReturnValueSurvivesSnapshot(t *testing.T) { + ps, _ := newPersistVM(t) + code := ` +local plugin = { name = "opinionated" } +plugin.state = { n = 0 } +plugin.hooks = { request_end = "tag" } +function plugin.tag(payload) + plugin.state.n = plugin.state.n + 1 + return { cost_usd = 1.25, verdict = "billed" } +end +return plugin +` + if err := ps.LoadSource("opinionated", code); err != nil { + t.Fatal(err) + } + out := ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "source": "s"}) + if out["cost_usd"] != 1.25 { + t.Errorf("★ hook return value lost: payload = %v", out) + } + if out["verdict"] != "billed" { + t.Errorf("merged field lost: %v", out) + } + ps.save.flush() + st := ps.State("opinionated").(map[string]interface{}) + if st["n"] != float64(1) { + t.Errorf("state.n = %v, want 1 (the hook still ran)", st["n"]) + } +} + +// TestCorruptStateFileIsNotFatal: a truncated write must degrade to compiled-in +// defaults, never to a gateway that refuses to start. +func TestCorruptStateFileIsNotFatal(t *testing.T) { + ps, _ := newPersistVM(t) + loadCounter(t, ps) + for i := 0; i < 3; i++ { + ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "source": "s"}) + } + ps.save.flush() + ps.Close() + + if err := os.WriteFile(ps.stateFile("counter"), []byte("{not json"), 0o644); err != nil { + t.Fatal(err) + } + + vm := NewVM(filepath.Join(t.TempDir(), "adapters")) + if err := vm.Start(); err != nil { + t.Fatal(err) + } + defer vm.Stop() + ps2 := NewPlugins(vm, ps.dir) + defer ps2.Close() + var warned bool + ps2.logf = func(string, ...interface{}) { warned = true } + loadCounter(t, ps2) + + if got := stateN(t, ps2); got != 0 { + t.Errorf("with a corrupt file the plugin must fall back to its defaults, got n = %v", got) + } + if !warned { + t.Error("a corrupt state file must warn the operator, not fail silently") + } + // And it must still forward. + ps2.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "source": "s"}) +} + +// TestFlushIsCoalesced: N mutations must not become N writes. A synchronous +// per-request write would put a file write on the hot path. +func TestFlushIsCoalesced(t *testing.T) { + ps, _ := newPersistVM(t) + loadCounter(t, ps) + for i := 0; i < 500; i++ { + ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "source": "s"}) + } + ps.save.flush() + // 500 mutations must collapse to a handful of writes, not 500. The bound is + // loose on purpose: the saver also flushes on a timer, so how many writes + // happen during the loop depends on how long 500 hook calls take. What must + // never happen is one write per mutation. + if w := ps.save.writes.Load(); w > 5 { + t.Errorf("500 hook calls produced %d file writes; the saver must coalesce", w) + } + if got := stateN(t, ps); got != 500 { + t.Fatalf("n = %v, want 500", got) + } +} + +// TestCloseFlushesTail: the shutdown path must not lose the last interval, +// which would reintroduce the same defect in a smaller window. +func TestCloseFlushesTail(t *testing.T) { + ps, _ := newPersistVM(t) + loadCounter(t, ps) + ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "source": "s"}) + ps.Close() + // Close() already waits for the final flush; no sleep needed. + + b, err := os.ReadFile(ps.stateFile("counter")) + if err != nil { + t.Fatalf("★ Close() did not flush: %v — a shutdown loses the tail", err) + } + var saved struct { + State map[string]interface{} `json:"state"` + } + if err := json.Unmarshal(b, &saved); err != nil { + t.Fatal(err) + } + if saved.State["n"] != float64(1) { + t.Errorf("flushed n = %v, want 1", saved.State["n"]) + } +} diff --git a/internal/lua/plugins.go b/internal/lua/plugins.go new file mode 100644 index 0000000..ce3f5c8 --- /dev/null +++ b/internal/lua/plugins.go @@ -0,0 +1,1584 @@ +package lua + +// Plugin runtime. +// +// A plugin is a single .lua file, loaded from its own directory, that extends +// the gateway in two ways: +// +// 1. Hooks: it registers callbacks on the request pipeline's stages +// (request_start, response_end, …). A hook receives a JSON table and +// returns either nil (no opinion) or a JSON object. +// 2. UI: at boot it returns HTML/CSS/JS fragments that the kernel injects +// into the WebUI — either as a whole new page, or as an extra element on +// an existing page. +// +// Design notes that are load-bearing (each one cost something to learn): +// +// - Plugins run in a SEPARATE VM from adapters, and each plugin gets its own +// elastic pool, exactly like an adapter. Sharing one state would let a +// plugin's globals corrupt an adapter's protocol translation (or vice +// versa), and a plugin is third-party code while an adapter is core. +// +// - A plugin that errors must NEVER break request forwarding. Hook calls are +// wrapped so a plugin error is logged and the original value is returned +// unchanged. A broken plugin is a missing feature, not an outage — the same +// reason transform_stream_chunk swallows errors today. +// +// - The billing plugin therefore cannot be trusted to be the source of truth +// for anything the gateway must enforce. It reads usage off the hook +// payload and accumulates in its own state; the gateway's own quota +// accounting (internal/gateway/stats.go) stays authoritative for limits. +// Two accounting paths that disagree is worse than one that is slightly +// less featureful, so the split is explicit and documented. + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "sync/atomic" + "time" + + golua "github.com/aarzilli/golua/lua" +) + +// Stage identifies a point in the request pipeline. Plugins may register a +// function for any stage; unknown stages are ignored at call time so an old +// plugin survives a gateway that grew new stages. +type Stage string + +// The pipeline stages. The order here is the order they fire in; it is the +// contract plugins are written against. +const ( + // StageRequestStart fires after the gateway has parsed and authorized a + // request but BEFORE any upstream slot is chosen. payload: + // stage, type ("chat"|"stream"|"image"), model (as requested), + // key (masked gateway key id), role, source (empty), stream (bool), + // messages_count, tools_count, ts (unix seconds). + StageRequestStart Stage = "request_start" + + // StageRouted fires once a (source, model) slot has been chosen and before + // the upstream call. payload adds: source, model (the resolved one), + // tier (AUTO tier, -1 for the direct path), stream. + StageRouted Stage = "routed" + + // StageChainStep fires ONCE PER STEP of an AUTO chain walk, and only on the + // AUTO path (a direct request has no chain and therefore emits nothing). + // + // This exists because StageRouted cannot express degradation: it fires once, + // after the walk, with the slot that finally won. "tier 1 was cooling so we + // dropped to tier 3" and "tier 1 served it" were indistinguishable. That + // distinction is the whole point of a priority chain, and it is what an + // operator debugging "why did my expensive model not get used" needs. + // + // payload: + // kind "tier_skip" | "slot_fail" | "tier_busy" | "selected" + // tier the AUTO tier this step belongs to (1 = highest priority) + // source / model set for slot_fail and selected + // reason human-readable cause, for tier_skip and tier_busy + // error the underlying error text, for slot_fail + // attempt 1-based slot attempt within this walk + // + // Ordering: every step precedes StageRouted, and the "selected" step is the + // last one. A plugin accumulating the walk therefore has the full picture + // by the time request_end arrives. + // + // These events are OBSERVATION ONLY — see the accounting note in + // docs/plugins.md: nothing here feeds back into routing, cooldown or quota. + StageChainStep Stage = "chain_step" + + // StageRequestEnd fires exactly once per request, after the client response + // has been produced (or after a failure was recorded). payload adds: + // source, model, ok, status, latency_ms, first_byte_ms, prompt_tokens, + // completion_tokens, cache_hit_tokens, cache_miss_tokens, image_count, + // error ("" when ok). + // + // This is the stage a billing plugin should read: it carries the final + // accounting for the request, including the upstream's own usage numbers. + StageRequestEnd Stage = "request_end" +) + +// AllStages is the firing order, used by the docs and by the hook listing. +var AllStages = []Stage{ + StageRequestStart, + StageChainStep, + StageRouted, + StageRequestEnd, +} + +// UIExtension is what a plugin contributes to the WebUI at boot. +type UIExtension struct { + // Page is a whole new sidebar entry + pane. Requires PageID and Title. + // The kernel renders Page's HTML into a pane whose id is "tab-"+PageID and + // adds a sidebar button with data-tab="". + Page *UIPage `json:"page,omitempty"` + // Pages is the multi-page form of the same thing. A plugin with several + // distinct screens (billing totals vs. the price rules that produced + // them) would otherwise have to cram both into one pane behind tabs, or + // smuggle the second one in as a hidden element. Both make the sidebar + // lie about what the plugin contributes. + // + // Page and Pages merge: a plugin may use either or both. + Pages []*UIPage `json:"pages,omitempty"` + // Elements are snippets injected into EXISTING pages, keyed by target page + // id (e.g. "status", "keys"). Order within a target is plugin load order. + Elements []UIElement `json:"elements,omitempty"` +} + +// UIPage is a plugin-provided page. +type UIPage struct { + PageID string `json:"page_id"` // kebab-case; becomes data-tab and #tab- + Title string `json:"title"` // sidebar label + Icon string `json:"icon"` // optional inline SVG or short glyph + Order int `json:"order"` // sidebar sort key (default 100) + // Mount is the page body. It may contain +]==], + }, + -- Two elements on the EXISTING status page: a headline tile and a + -- per-source cost breakdown, so the number is visible without opening the + -- Billing tab. + elements = { + { + target = "status", + anchor = "top", + order = 5, + mount = [==[ +
+
+
—
+
+
+ +]==], + }, + }, +} + + +return plugin \ No newline at end of file diff --git a/internal/lua/plugins_test.go b/internal/lua/plugins_test.go new file mode 100644 index 0000000..5d3d26b --- /dev/null +++ b/internal/lua/plugins_test.go @@ -0,0 +1,377 @@ +package lua + +import ( + "encoding/json" + "os" + "path/filepath" + "strings" + "testing" +) + +// newPluginVM builds a VM plus a plugin registry rooted at dir. +func newPluginVM(t *testing.T) (*VM, *Plugins, string) { + t.Helper() + dir := filepath.Join(t.TempDir(), "adapters") + vm := NewVM(dir) + if err := vm.Start(); err != nil { + t.Fatalf("vm start: %v", err) + } + t.Cleanup(vm.Stop) + pdir := filepath.Join(t.TempDir(), "plugins") + return vm, NewPlugins(vm, pdir), pdir +} + +// loadPlugin writes one plugin to disk and loads it. +func loadPlugin(t *testing.T, ps *Plugins, name, code string) error { + t.Helper() + if err := os.MkdirAll(ps.dir, 0755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(ps.dir, name+".lua"), []byte(code), 0644); err != nil { + t.Fatal(err) + } + return ps.LoadSource(name, code) +} + +// TestPluginManifestAndHooks: the two hook registration forms both work and the +// manifest is read. +func TestPluginManifestAndHooks(t *testing.T) { + _, ps, _ := newPluginVM(t) + code := ` +local p = {} +p.name = "demo" +p.version = "1.2.3" +p.description = "a demo plugin" +p.author = "tester" +p.hooks = { request_end = "on_end" } +function p.on_end(payload) + payload.seen = true + payload.name_seen = "demo" + return payload +end +return p +` + if err := loadPlugin(t, ps, "demo", code); err != nil { + t.Fatalf("load: %v", err) + } + list := ps.List() + if len(list) != 1 { + t.Fatalf("List() = %d plugins, want 1", len(list)) + } + if list[0]["name"] != "demo" || list[0]["version"] != "1.2.3" { + t.Errorf("manifest not read: %+v", list[0]) + } + hooks := list[0]["hooks"].([]string) + if len(hooks) != 1 || hooks[0] != string(StageRequestEnd) { + t.Errorf("hooks = %v, want [request_end]", hooks) + } + out := ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m"}) + if out["seen"] != true || out["name_seen"] != "demo" { + t.Errorf("hook did not mutate payload: %+v", out) + } +} + +// TestPluginAnonymousHookForm: `request_end = function() end` directly on the +// table must register too, since a single-hook plugin should not need a name. +func TestPluginAnonymousHookForm(t *testing.T) { + _, ps, _ := newPluginVM(t) + code := ` +local p = { name = "anon" } +p.request_end = function(payload) + payload.hit = 1 + return payload +end +return p +` + if err := loadPlugin(t, ps, "anon", code); err != nil { + t.Fatalf("load: %v", err) + } + out := ps.Fire(StageRequestEnd, map[string]interface{}{}) + if out["hit"] != float64(1) { + t.Errorf("anonymous hook did not fire: %+v", out) + } +} + +// TestPluginHookStagesFireInOrder: each stage reaches only its own hooks. +func TestPluginHookStagesFireInOrder(t *testing.T) { + _, ps, _ := newPluginVM(t) + code := ` +local p = { name = "stages" } +p.hooks = { + request_start = "s1", + routed = "s2", + request_end = "s3", +} +function p.s1(x) x.order = (x.order or "") .. "1" return x end +function p.s2(x) x.order = (x.order or "") .. "2" return x end +function p.s3(x) x.order = (x.order or "") .. "3" return x end +return p +` + if err := loadPlugin(t, ps, "stages", code); err != nil { + t.Fatalf("load: %v", err) + } + payload := map[string]interface{}{} + ps.Fire(StageRequestStart, payload) + ps.Fire(StageRouted, payload) + ps.Fire(StageRequestEnd, payload) + if payload["order"] != "123" { + t.Errorf("stage order = %v, want \"123\"", payload["order"]) + } +} + +// TestPluginErrorIsContained is the critical safety property: a throwing hook +// must not propagate. Forwarding depends on it. +func TestPluginErrorIsContained(t *testing.T) { + _, ps, _ := newPluginVM(t) + code := ` +local p = { name = "boom" } +p.hooks = { request_end = "kaboom" } +function p.kaboom(payload) + error("intentional plugin failure") +end +return p +` + if err := loadPlugin(t, ps, "boom", code); err != nil { + t.Fatalf("load: %v", err) + } + // Must not panic and must return the payload unchanged. + out := ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m"}) + if out["model"] != "m" { + t.Errorf("payload was altered by a failing plugin: %+v", out) + } + // And the failure must be visible, not silent. + errs := ps.HookErrors() + if errs["request_end"] == nil { + t.Error("a failing plugin left no error record; it would be silently missing") + } +} + +// TestPluginFailingHookDoesNotBlockLaterPlugins: one bad plugin must not stop +// the next one from running. +func TestPluginFailingHookDoesNotBlockLaterPlugins(t *testing.T) { + _, ps, _ := newPluginVM(t) + bad := ` +local p = { name = "bad" } +p.hooks = { request_end = "f" } +function p.f(x) error("boom") end +return p +` + good := ` +local p = { name = "good" } +p.hooks = { request_end = "f" } +function p.f(x) x.good = true return x end +return p +` + _ = loadPlugin(t, ps, "bad", bad) + if err := loadPlugin(t, ps, "good", good); err != nil { + t.Fatalf("load good: %v", err) + } + out := ps.Fire(StageRequestEnd, map[string]interface{}{}) + if out["good"] != true { + t.Errorf("a good plugin was blocked by a failing one: %+v", out) + } +} + +// TestPluginSyntaxErrorIsIsolated: a plugin that will not compile is listed +// with its error and is never called — it must not prevent LoadDir from loading +// the rest. +func TestPluginSyntaxErrorIsIsolated(t *testing.T) { + _, ps, _ := newPluginVM(t) + broken := "this is not lua(((" + good := ` +local p = { name = "ok" } +p.hooks = { request_end = "f" } +function p.f(x) x.ok = true return x end +return p +` + _ = loadPlugin(t, ps, "broken", broken) + if err := loadPlugin(t, ps, "ok", good); err != nil { + t.Fatalf("load ok: %v", err) + } + if err := ps.LoadDir(); err != nil { + t.Fatalf("LoadDir: %v", err) + } + // The broken plugin must not be callable and must carry an error. + for _, row := range ps.List() { + if row["name"] == "broken" { + if row["loaded"] == true { + t.Error("a plugin with a syntax error reported itself as loaded") + } + if row["error"] == nil || row["error"] == "" { + t.Error("a broken plugin carries no error message") + } + } + } + // The good plugin still works. + out := ps.Fire(StageRequestEnd, map[string]interface{}{}) + if out["ok"] != true { + t.Errorf("good plugin stopped working: %+v", out) + } +} + +// TestPluginUIExtension: a plugin can contribute a page and elements. +func TestPluginUIExtension(t *testing.T) { + _, ps, _ := newPluginVM(t) + code := ` +local p = { name = "ui" } +p.hooks = { request_end = "f" } +function p.f(x) return x end +p.ui = { + page = { + page_id = "billing", + title = "Billing", + icon = "💰", + order = 50, + mount = "
hi
", + }, + elements = { + { target = "status", anchor = "top", mount = "
cost
" }, + }, +} +return p +` + if err := loadPlugin(t, ps, "ui", code); err != nil { + t.Fatalf("load: %v", err) + } + // Every contributed page — including the single `page` field — is folded + // into one merged list. Reading ui.Page here would silently pass on a + // payload whose pages were all dropped. + ui := ps.UI() + if len(ui.Pages) != 1 { + t.Fatalf("expected 1 merged page, got %d", len(ui.Pages)) + } + pg := ui.Pages[0] + if pg.PageID != "billing" || pg.Title != "Billing" { + t.Errorf("page = %+v", pg) + } + if !strings.Contains(pg.Mount, "console.log") { + t.Error("mount lost its script content") + } + if len(ui.Elements) != 1 || ui.Elements[0].Target != "status" { + t.Errorf("elements = %+v", ui.Elements) + } +} + +// TestPluginHookReturnsNilIsNoOpinion: a hook returning nothing must leave the +// payload untouched (plugins should not be forced to echo it back). +func TestPluginHookReturnsNilIsNoOpinion(t *testing.T) { + _, ps, _ := newPluginVM(t) + code := ` +local p = { name = "silent" } +p.hooks = { request_end = "f" } +function p.f(payload) + -- records nothing, returns nothing + return nil +end +return p +` + if err := loadPlugin(t, ps, "silent", code); err != nil { + t.Fatalf("load: %v", err) + } + out := ps.Fire(StageRequestEnd, map[string]interface{}{"model": "m", "ok": true}) + if out["model"] != "m" || out["ok"] != true { + t.Errorf("a no-op hook disturbed the payload: %+v", out) + } +} + +// TestPluginUIJSONShape is a wire-format guard: the kernel sends this to the +// browser, so the shape is a contract with the WebUI. +func TestPluginUIJSONShape(t *testing.T) { + _, ps, _ := newPluginVM(t) + code := ` +local p = { name = "shape" } +p.hooks = { request_end = "f" } +function p.f(x) return x end +p.ui = { elements = { { target = "keys", mount = "k" } } } +return p +` + if err := loadPlugin(t, ps, "shape", code); err != nil { + t.Fatalf("load: %v", err) + } + b, err := json.Marshal(ps.UI()) + if err != nil { + t.Fatalf("marshal UI: %v", err) + } + var view struct { + Elements []struct { + Target string `json:"target"` + Mount string `json:"mount"` + } `json:"elements"` + } + if err := json.Unmarshal(b, &view); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if len(view.Elements) != 1 || view.Elements[0].Target != "keys" { + t.Errorf("UI JSON shape = %+v", view.Elements) + } +} + +// TestPluginFireWithNoPluginsIsNoop: an empty registry must not allocate or fail. +func TestPluginFireWithNoPluginsIsNoop(t *testing.T) { + _, ps, _ := newPluginVM(t) + in := map[string]interface{}{"a": 1} + out := ps.Fire(StageRequestEnd, in) + if out["a"] != 1 || ps.Count() != 0 { + t.Errorf("empty registry misbehaved: %+v", out) + } +} + +// TestChainStepIsARealStage: chain_step is documented as a distinct stage that +// fires once per step of an AUTO walk. If it were only a field on `routed`, a +// plugin author following the docs would silently get one event instead of the +// whole walk. +func TestChainStepIsARealStage(t *testing.T) { + found := false + for _, s := range AllStages { + if s == StageChainStep { + found = true + } + } + if !found { + t.Fatal("StageChainStep is not in AllStages, so the dispatcher never registers it") + } + // Ordering: it must sit between request_start and routed, which is what + // docs/plugins.md promises. + var iStart, iStep, iRouted = -1, -1, -1 + for i, s := range AllStages { + switch s { + case StageRequestStart: + iStart = i + case StageChainStep: + iStep = i + case StageRouted: + iRouted = i + } + } + if !(iStart < iStep && iStep < iRouted) { + t.Errorf("stage order = %v, want request_start < chain_step < routed", AllStages) + } + _, ps, _ := newPluginVM(t) + code := ` +local p = { name = "stepper" } +p.hooks = { chain_step = "s" } +function p.s(payload) + payload.kinds = (payload.kinds or "") + return nil +end +return p +` + if err := loadPlugin(t, ps, "stepper", code); err != nil { + t.Fatalf("load: %v", err) + } + // It must actually dispatch. + seen := false + for _, row := range ps.List() { + if row["name"] == "stepper" { + hooks := row["hooks"].([]string) + for _, h := range hooks { + if h == string(StageChainStep) { + seen = true + } + } + } + } + if !seen { + t.Error("a plugin registered for chain_step is not reported as such") + } +} diff --git a/internal/lua/vm.go b/internal/lua/vm.go index d22e7bc..9ad0e45 100644 --- a/internal/lua/vm.go +++ b/internal/lua/vm.go @@ -36,7 +36,7 @@ import ( golua "github.com/aarzilli/golua/lua" ) -//go:embed adapters/*.lua +//go:embed adapters/*.lua plugins/*.lua var bundledAdapters embed.FS // adapterGlobal is the reserved global holding the adapter table after the @@ -103,6 +103,11 @@ type adapterPool struct { lastGrow time.Time idleRounds int // consecutive janitor rounds that saw reclaimable slack peakInUse int // high-water mark of inUse, for observability + // pluginMode makes boot() store the returned table under pluginGlobal + // instead of adapterGlobal. Everything else (elastic sizing, reclaim) is + // identical, which is why plugins reuse this pool rather than getting a + // second implementation. + pluginMode bool } const ( @@ -116,6 +121,13 @@ const ( // growCooldown keeps a burst of misses from batching repeatedly while the // previous batch is still booting. growCooldown = time.Second + // maxPluginStates caps how many concurrent VM states ONE plugin may occupy. + // A plugin is third-party code on the request path, so its ceiling is much + // lower than an adapter's (which is sized from the sources' + // max_concurrent): a plugin hook is a short synchronous call, so a handful + // of states is already far more parallelism than any real hook needs, and a + // runaway plugin cannot balloon memory the way a per-source adapter pool can. + maxPluginStates = 4 // residentWorkers is how many states an adapter keeps warm once it has // served at least one request. Booting is milliseconds, but keeping one warm // removes that from the critical path of the next request. Adapters that @@ -265,7 +277,14 @@ func (p *adapterPool) boot() (*worker, error) { L.Close() return nil, fmt.Errorf("adapter %s must return a table", p.name) } - L.SetGlobal(adapterGlobal) + // SetGlobal POPS the value off the stack, so it can only be called once per + // boot. Plugins therefore store under pluginGlobal only, and NOT under + // adapterGlobal: a second SetGlobal on the now-empty stack would assign nil. + if p.pluginMode { + L.SetGlobal(pluginGlobal) + } else { + L.SetGlobal(adapterGlobal) + } L.SetTop(0) return &worker{L: L}, nil } @@ -601,6 +620,55 @@ func (v *VM) Start() error { return nil } +// bundledPluginsDir is the directory inside the embedded FS holding the +// plugins shipped with the gateway. It is separate from adapters/ on purpose: +// the two are loaded by different machinery into different kinds of Lua state +// (a protocol transform vs. request-pipeline hooks), and keeping them apart +// makes it obvious that dropping a file in one does not affect the other. +const bundledPluginsDir = "plugins" + +// ReadBundledPlugin returns the source of a plugin shipped with the gateway. +// It exists so a test (or an operator tool) can load a bundled plugin without +// depending on whether seeding has already run for this directory. +func ReadBundledPlugin(name string) (string, error) { + data, err := bundledAdapters.ReadFile(bundledPluginsDir + "/" + name + ".lua") + if err != nil { + return "", fmt.Errorf("bundled plugin %s: %w", name, err) + } + return string(data), nil +} + +// writeBundledPlugins seeds the plugin directory with the shipped plugins. +// +// It runs only when the directory does not exist yet (same rule as adapters): +// once the directory exists it is authoritative, so deleting a shipped plugin is +// a real delete and editing one survives restarts. +func writeBundledPlugins(dir string) error { + if dir == "" { + return nil + } + if err := os.MkdirAll(dir, 0755); err != nil { + return fmt.Errorf("mkdir plugin dir: %w", err) + } + entries, err := bundledAdapters.ReadDir(bundledPluginsDir) + if err != nil { + return nil // nothing embedded; not an error + } + for _, e := range entries { + if e.IsDir() || filepath.Ext(e.Name()) != ".lua" { + continue + } + data, err := bundledAdapters.ReadFile(bundledPluginsDir + "/" + e.Name()) + if err != nil { + continue + } + if err := os.WriteFile(filepath.Join(dir, e.Name()), data, 0644); err != nil { + return err + } + } + return nil +} + func (v *VM) Stop() { select { case <-v.janitorStop: @@ -925,9 +993,47 @@ func restorePcall(L *golua.State) { } } +// hookGuardName is the Lua global that wraps every plugin hook call. +const hookGuardName = "__llmsproxy_call_hook" + +// hookGuardSrc defines the wrapper. It exists because of a hard constraint in +// the binding: golua's callEx calls L.StackTrace() on ANY pcall error, and +// StackTrace() calls lua_getinfo, which SIGSEGVs in this LuaJIT build once the +// stack is deep enough (a request_end payload with the AUTO chain trace does +// it). Go cannot recover from a C-level signal, so a single Lua mistake inside +// one plugin killed the whole gateway — the production outage this was written +// for. Repeated crashes proved the boundary is not theoretical. +// +// pcall INSIDE Lua catches the error before golua ever sees a non-zero +// pcall status, so the C stack-trace path is never entered. The failure comes +// back as an ordinary (nil, message) pair, which the Go side records in +// hook_errors and moves on from — the documented contract that "a broken plugin +// must not affect request forwarding" finally holds for script errors too, not +// just for Go panics. +// +// returns: (result, errorMessage) — both nil/"" on success. +const hookGuardSrc = ` +function ` + hookGuardName + `(fn, payload) + local ok, res = pcall(fn, payload) + if not ok then + return nil, tostring(res) + end + return res, nil +end` + +func registerHookGuard(L *golua.State) { + if err := L.DoString(hookGuardSrc); err != nil { + // Nothing useful to do here beyond leaving the global absent: invoke() + // checks for it and falls back to a direct call, which still works for + // correct plugins (only their ERRORS stop being survivable). + L.Pop(1) + } +} + func setupGlobals(L *golua.State) { restorePcall(L) buildJSONTable(L) + registerHookGuard(L) registerFn(L, "hmac_sha256_hex", func(L *golua.State) int { key := L.ToString(1) @@ -1104,6 +1210,56 @@ func pushGoValue(L *golua.State, v interface{}) { } func jsonEncode(v interface{}) ([]byte, error) { return json.Marshal(v) } + +// luaToJSON converts the Lua value at idx into a Go value via json.encode, then +// unmarshals it into out. It is the bridge used by the plugin manifest/UI +// reader: the plugin returns a plain Lua table, and Go wants a typed struct. +// +// It goes through JSON rather than walking the Lua stack directly because the +// adapter/plugin boundary already speaks JSON everywhere else (transform_request +// gets a JSON string, hooks get a JSON string), so this keeps one representation +// instead of two. +func luaToJSON(L *golua.State, idx int, out interface{}) error { + if L.GetTop() < 1 { + return fmt.Errorf("empty stack") + } + abs := idx + if abs < 0 { + abs = L.GetTop() + 1 + abs + } + if abs < 1 || abs > L.GetTop() { + return fmt.Errorf("index %d out of range (top=%d)", idx, L.GetTop()) + } + // Absolute indices throughout: this binding aborts the process (SIGABRT) + // on a bad index rather than panicking, so the stack is captured before any + // push instead of being addressed relative to a shifting top. + // + // json.encode is pushed onto the stack and the value is pushed AFTER it, so + // Call(1, 1) invokes it (Call takes no function index — it calls whatever + // sits below the nargs values). + L.GetGlobal("json") + if L.IsNil(-1) { + L.SetTop(0) + return fmt.Errorf("json global missing") + } + L.GetField(-1, "encode") + if L.Type(-1) != golua.LUA_TFUNCTION { + L.SetTop(0) + return fmt.Errorf("json.encode missing") + } + L.PushValue(abs) + if err := L.Call(1, 1); err != nil { + L.SetTop(0) + return err + } + if L.GetTop() < 1 || L.Type(-1) != golua.LUA_TSTRING { + L.SetTop(0) + return fmt.Errorf("json.encode did not return a string") + } + s := L.ToString(-1) + L.SetTop(0) + return json.Unmarshal([]byte(s), out) +} func jsonDecode(s string) (interface{}, error) { var v interface{} if err := json.Unmarshal([]byte(s), &v); err != nil { diff --git a/internal/provider/gemini_endpoint_test.go b/internal/provider/gemini_endpoint_test.go new file mode 100644 index 0000000..7267efc --- /dev/null +++ b/internal/provider/gemini_endpoint_test.go @@ -0,0 +1,270 @@ +package provider + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "sync" + "testing" + + "llmsproxy/internal/config" + "llmsproxy/internal/lua" + "llmsproxy/internal/types" +) + +// Gemini is the only adapter whose endpoint carries the model in the PATH +// (POST /v1beta/models/{model}:generateContent), so it is the only one that +// needs the substitution this file covers. +// +// The defect it fixes was silent: the adapter shipped endpoint="/v1/models", +// which is Gemini's model-LIST endpoint and answers 405 to POST, so the preset +// template produced a source that could never work. + +func geminiVM(t *testing.T) *lua.VM { + t.Helper() + // The dir must be named "adapters": the VM seeds its bundled adapters into + // on Start and a differently-named dir leaves it empty, which surfaces + // as "adapter gemini not loaded" rather than an obvious setup error. + vm := lua.NewVM(filepath.Join(t.TempDir(), "adapters")) + if err := vm.Start(); err != nil { + t.Fatalf("lua vm: %v", err) + } + t.Cleanup(vm.Stop) + return vm +} + +// geminiProvider builds a provider on the gemini adapter against base. +func geminiProvider(t *testing.T, base string, models ...string) *Provider { + t.Helper() + if len(models) == 0 { + models = []string{"gemini-3.6-flash"} + } + ms := make([]config.Model, 0, len(models)) + for _, m := range models { + ms = append(ms, config.Model{ID: m, Kind: "chat"}) + } + return New(config.Source{ + Name: "gem", + BaseURL: base, + Adapter: "gemini", + APIKey: "k", + MaxConcurrent: 4, + Models: ms, + }, geminiVM(t)) +} + +// geminiRecorder captures the path the gateway actually requested and replies +// with a minimal Gemini-shaped payload. +type geminiRecorder struct { + mu sync.Mutex + paths []string + bodies []string + stream bool +} + +func (g *geminiRecorder) handler(t *testing.T) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + raw, _ := io.ReadAll(r.Body) + g.mu.Lock() + // RequestURI, not URL.Path: URL.Path is the DECODED path, so an escaped + // %2F looks like a real "/" and the assertion could not tell an escaped + // id from an unescaped one. + g.paths = append(g.paths, r.RequestURI) + g.bodies = append(g.bodies, string(raw)) + stream := g.stream + g.mu.Unlock() + + if stream { + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(200) + fmt.Fprintf(w, "data: %s\n\n", `{"candidates":[{"content":{"parts":[{"text":"hi"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":3,"candidatesTokenCount":1,"totalTokenCount":4}}`) + fmt.Fprintln(w, "data: [DONE]") + return + } + w.Header().Set("Content-Type", "application/json") + fmt.Fprint(w, `{"candidates":[{"content":{"parts":[{"text":"pong"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":3,"candidatesTokenCount":1,"totalTokenCount":4}}`) + } +} + +func (g *geminiRecorder) lastPath() string { + g.mu.Lock() + defer g.mu.Unlock() + if len(g.paths) == 0 { + return "" + } + return g.paths[len(g.paths)-1] +} + +// TestGeminiEndpointCarriesModel: a non-streaming chat must POST to +// /v1beta/models/:generateContent. The old static endpoint produced +// "/v1beta/v1/models" (a list endpoint, 405 on POST). +func TestGeminiEndpointCarriesModel(t *testing.T) { + rec := &geminiRecorder{} + srv := httptest.NewServer(rec.handler(t)) + defer srv.Close() + + p := geminiProvider(t, srv.URL, "gemini-3.6-flash") + if _, err := p.Chat(context.Background(), &types.ChatRequest{ + Model: "gemini-3.6-flash", + Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}}, + }); err != nil { + t.Fatalf("chat: %v (path was %q)", err, rec.lastPath()) + } + if got, want := rec.lastPath(), "/v1beta/models/gemini-3.6-flash:generateContent"; got != want { + t.Errorf("path = %q, want %q", got, want) + } +} + +// TestGeminiStreamUsesStreamGenerateContent: the streaming verb lives in the +// path too, so it must be swapped for stream requests only. +func TestGeminiStreamUsesStreamGenerateContent(t *testing.T) { + rec := &geminiRecorder{stream: true} + srv := httptest.NewServer(rec.handler(t)) + defer srv.Close() + + p := geminiProvider(t, srv.URL, "gemini-3.6-flash") + ch, err := p.ChatStream(context.Background(), &types.ChatRequest{ + Model: "gemini-3.6-flash", + Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}}, + }) + if err != nil { + t.Fatalf("stream: %v (path was %q)", err, rec.lastPath()) + } + for range ch { + } + if got, want := rec.lastPath(), "/v1beta/models/gemini-3.6-flash:streamGenerateContent"; got != want { + t.Errorf("path = %q, want %q", got, want) + } +} + +// TestGeminiModelFollowsRequestNotSourceDefault is the part a naive fix gets +// wrong: substituting the source's default model would make every request on a +// multi-model source bill the wrong model. AUTO pins req.Model per slot, so the +// URL must follow the REQUEST's model. +func TestGeminiModelFollowsRequestNotSourceDefault(t *testing.T) { + rec := &geminiRecorder{} + srv := httptest.NewServer(rec.handler(t)) + defer srv.Close() + + p := geminiProvider(t, srv.URL, "gemini-3.6-flash", "gemini-3.1-pro") + for _, m := range []string{"gemini-3.1-pro", "gemini-3.6-flash"} { + if _, err := p.Chat(context.Background(), &types.ChatRequest{ + Model: m, + Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}}, + }); err != nil { + t.Fatalf("chat %s: %v", m, err) + } + if want := "/v1beta/models/" + m + ":generateContent"; rec.lastPath() != want { + t.Errorf("path = %q, want %q", rec.lastPath(), want) + } + } +} + +// TestGeminiModelIsEscaped: a model id reaches a URL path, so an id carrying a +// slash or space must be escaped rather than silently addressing a different +// resource (or producing a request Go's http client rejects). +func TestGeminiModelIsEscaped(t *testing.T) { + rec := &geminiRecorder{} + srv := httptest.NewServer(rec.handler(t)) + defer srv.Close() + + // Configure the odd id as a real model so ModelFor resolves it. + p := geminiProvider(t, srv.URL, "vendor/model:v1 beta") + if _, err := p.Chat(context.Background(), &types.ChatRequest{ + Model: "vendor/model:v1 beta", + Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}}, + }); err != nil { + t.Fatalf("chat: %v", err) + } + got := rec.lastPath() + if strings.Contains(got, "/v1beta/models/vendor/model") { + t.Errorf("path %q was not escaped: the id's slash changed the resource", got) + } + if !strings.HasPrefix(got, "/v1beta/models/") || !strings.HasSuffix(got, ":generateContent") { + t.Errorf("path = %q, want an escaped id between the prefix and the verb", got) + } +} + +// TestNonGeminiEndpointsAreUntouched: the substitution and the verb rewrite must +// be no-ops for every adapter whose endpoint is a fixed path. A regression here +// would silently rewrite OpenAI/Anthropic/… URLs. +func TestNonGeminiEndpointsAreUntouched(t *testing.T) { + for _, adapter := range []string{"openai", "deepseek", "anthropic", "ollama", "mistral", "trae", "sensenova"} { + vm := geminiVM(t) + p := New(config.Source{ + Name: "s", BaseURL: "https://example.test", Adapter: adapter, + APIKey: "k", MaxConcurrent: 2, + Models: []config.Model{{ID: "m1", Kind: "chat"}}, + }, vm) + for _, stream := range []bool{false, true} { + got := p.ChatURL("m1", stream) + want := "https://example.test" + p.Endpoint() + if got != want { + t.Errorf("adapter %s stream=%v: URL = %q, want %q", adapter, stream, got, want) + } + if strings.Contains(got, "streamGenerateContent") { + t.Errorf("adapter %s: streaming verb leaked into a fixed endpoint: %q", adapter, got) + } + } + vm.Stop() + } +} + +// TestSourceEndpointOverridesTemplate: a source that spells its own endpoint +// must win outright — including when it spells one WITHOUT the placeholder. +func TestSourceEndpointOverridesTemplate(t *testing.T) { + vm := geminiVM(t) + defer vm.Stop() + p := New(config.Source{ + Name: "s", BaseURL: "https://example.test", Adapter: "gemini", + Endpoint: "/custom/path", APIKey: "k", MaxConcurrent: 2, + Models: []config.Model{{ID: "m1", Kind: "chat"}}, + }, vm) + if got, want := p.ChatURL("m1", false), "https://example.test/custom/path"; got != want { + t.Errorf("URL = %q, want %q (source endpoint must override the adapter template)", got, want) + } + // And the streaming verb rewrite must not damage an unrelated path. + if got := p.ChatURL("m1", true); strings.Contains(got, "streamGenerate") { + t.Errorf("streaming rewrite leaked into a user-spelled path: %q", got) + } +} + +// TestGeminiResponseStillParses: the endpoint fix must not disturb the body +// transform — a request that reaches the right URL still has to come back as a +// unified response with Gemini's usageMetadata mapped. +func TestGeminiResponseStillParses(t *testing.T) { + rec := &geminiRecorder{} + srv := httptest.NewServer(rec.handler(t)) + defer srv.Close() + + p := geminiProvider(t, srv.URL, "gemini-3.6-flash") + resp, err := p.Chat(context.Background(), &types.ChatRequest{ + Model: "gemini-3.6-flash", + Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}}, + }) + if err != nil { + t.Fatalf("chat: %v", err) + } + if resp.Content != "pong" { + t.Errorf("content = %q, want %q", resp.Content, "pong") + } + if resp.TokenUsage.Prompt != 3 || resp.TokenUsage.Completion != 1 || resp.TokenUsage.Total != 4 { + t.Errorf("usage = %+v, want prompt 3 / completion 1 / total 4", resp.TokenUsage) + } + // The body the adapter produced must still be Gemini-native. + rec.mu.Lock() + body := rec.bodies[0] + rec.mu.Unlock() + var sent map[string]interface{} + if err := json.Unmarshal([]byte(body), &sent); err != nil { + t.Fatalf("adapter sent invalid json: %v (%s)", err, body) + } + if _, ok := sent["contents"]; !ok { + t.Errorf("adapter did not convert to Gemini's contents[]: %s", body) + } +} diff --git a/internal/provider/provider.go b/internal/provider/provider.go index eb1564b..255e64e 100644 --- a/internal/provider/provider.go +++ b/internal/provider/provider.go @@ -531,7 +531,9 @@ func isAutoID(s string) bool { return s == "" || strings.EqualFold(s, "AUTO") } -// Endpoint resolves the upstream chat path. +// Endpoint resolves the upstream chat path. It may contain the placeholder +// "{model}"; use URL (or ChatURL) rather than calling this directly when the +// path has to be usable. func (p *Provider) Endpoint() string { if p.cfg.Endpoint != "" { return p.cfg.Endpoint @@ -553,8 +555,40 @@ func (p *Provider) ImageEndpoint() string { return "/v1/images/generations" } +// modelPlaceholder is the marker an adapter puts in its endpoint template when +// the upstream carries the model id in the PATH rather than in the body. +// Gemini is the only such adapter: POST /v1beta/models/{model}:generateContent. +// Every other adapter's endpoint is a fixed path, so substituting is a no-op +// for them. +const modelPlaceholder = "{model}" + +// ChatURL resolves the full chat URL for a specific model. +// +// Two substitutions happen here, both driven by the adapter's endpoint template: +// +// - "{model}" is replaced by the model id actually being sent. Without this +// the gateway would POST to Gemini's model-LIST endpoint, which answers 405. +// - for a streaming request, the ":generateContent" verb becomes +// ":streamGenerateContent". Gemini streams over the same path with a +// different verb, and the verb is part of the path, so the two cannot both +// be static. The rewrite is deliberately narrow: it only fires on the exact +// ":generateContent" suffix, so an adapter whose endpoint merely mentions +// the word keeps its path untouched. +func (p *Provider) ChatURL(model string, stream bool) string { + ep := p.Endpoint() + if strings.Contains(ep, modelPlaceholder) { + // A model id is put in a URL path, so it must be escaped: an id with a + // slash would otherwise silently address a different resource. + ep = strings.ReplaceAll(ep, modelPlaceholder, url.PathEscape(model)) + } + if stream { + ep = strings.Replace(ep, ":generateContent", ":streamGenerateContent", 1) + } + return strings.TrimRight(p.cfg.BaseURL, "/") + ep +} + func (p *Provider) URL() string { - return strings.TrimRight(p.cfg.BaseURL, "/") + p.Endpoint() + return p.ChatURL("", false) } func (p *Provider) ImageURL() string { @@ -750,10 +784,11 @@ func (p *Provider) probeChat(ctx context.Context) (bool, string) { body, err := json.Marshal(probe) if err == nil { var hdr http.Header - if hdrs, herr := p.buildHeaders(string(body), p.URL(), ""); herr == nil { + probeURL := p.ChatURL(model, false) + if hdrs, herr := p.buildHeaders(string(body), probeURL, ""); herr == nil { hdr = hdrs } - raw, status, derr := p.do(ctx, p.URL(), string(body), hdr) + raw, status, derr := p.do(ctx, probeURL, string(body), hdr) if derr != nil { msg = derr.Error() } else if status == 200 { @@ -1084,11 +1119,12 @@ func (p *Provider) Chat(ctx context.Context, req *types.ChatRequest) (*types.Uni if err != nil { return nil, err } - hdrs, err := p.buildHeaders(body, p.URL(), req.ClientSession) + chatURL := p.ChatURL(model, false) + hdrs, err := p.buildHeaders(body, chatURL, req.ClientSession) if err != nil { return nil, err } - raw, status, err := p.do(ctx, p.URL(), body, hdrs) + raw, status, err := p.do(ctx, chatURL, body, hdrs) if err != nil { // a client disconnect or cancelled context is neither a success nor // a failure for scheduling purposes — only upstream errors count @@ -1144,7 +1180,10 @@ func (p *Provider) ChatStream(ctx context.Context, req *types.ChatRequest) (<-ch p.Release() return nil, err } - hdrs, err := p.buildHeaders(body, p.URL(), req.ClientSession) + // stream=true so a path-carried model plus the streaming verb is resolved + // for THIS request's model, not the source default. + streamURL := p.ChatURL(model, true) + hdrs, err := p.buildHeaders(body, streamURL, req.ClientSession) if err != nil { p.Release() return nil, err @@ -1156,7 +1195,7 @@ func (p *Provider) ChatStream(ctx context.Context, req *types.ChatRequest) (<-ch } rc := make(chan respOrErr, 1) go func() { - resp, err := p.doRawStream(ctx, p.URL(), body, hdrs) + resp, err := p.doRawStream(ctx, streamURL, body, hdrs) rc <- respOrErr{resp, err} }() diff --git a/internal/scheduler/empty_response_test.go b/internal/scheduler/empty_response_test.go new file mode 100644 index 0000000..b86369d --- /dev/null +++ b/internal/scheduler/empty_response_test.go @@ -0,0 +1,127 @@ +package scheduler + +import ( + "context" + "encoding/json" + "strings" + "testing" + "time" + + "llmsproxy/internal/types" +) + +// TestResultIsEmptyMatchesTheReasoningOnlyResponse guards the production bug +// behind "model AUTO returned a completed response with no content". +// +// Production measurement: 2 of 20 AUTO requests returned zero content, both +// served by claude-opus-4-8. Direct upstream capture showed 28 chunks of +// `reasoning_content` and NO `content`, finish_reason=length — the token budget +// was consumed by thinking before any text was emitted. runTier accepted it +// because err == nil. +func TestResultIsEmptyMatchesTheReasoningOnlyResponse(t *testing.T) { + // Exactly what the upstream returned, as a decoded response. + onlyReasoning := &types.UnifiedResponse{ + Model: "claude-opus-4-8", + ReasoningContent: "Alright, the user just said \"hi\". Simple greeting…", + FinishReason: "length", + TokenUsage: types.TokenUsage{Prompt: 14, Completion: 30, Total: 44}, + } + if !resultIsEmpty(onlyReasoning) { + t.Error("a reasoning-only response with usage must count as empty; " + + "it is what the client sees as \"completed with no content\"") + } + + // The same response WITH text is a normal answer. + if resultIsEmpty(&types.UnifiedResponse{Content: "Hi! 👋", ReasoningContent: "hmm"}) { + t.Error("a response with content is not empty") + } + // A tool-calling agent turn has no text and is perfectly valid. + if resultIsEmpty(&types.UnifiedResponse{ + ToolCalls: []types.ToolCall{{Name: "read_file", Arguments: map[string]interface{}{"path": "x"}}}, + }) { + t.Error("a tool-call-only response must NOT count as empty — agents " + + "legitimately produce tool calls with no text") + } + // An image slot returns no text by design. + if resultIsEmpty(&types.UnifiedResponse{ + ImageData: []types.ImageData{{URL: "http://x/y.png"}}, + }) { + t.Error("an image response must NOT count as empty") + } + // Whitespace-only content is as useless to a client as none at all. + if !resultIsEmpty(&types.UnifiedResponse{Content: " \n\t"}) { + t.Error("whitespace-only content must count as empty") + } + if resultIsEmpty(nil) { + t.Error("nil is not an empty response; it is an absent one") + } +} + +// TestPeekStreamHoldsReasoningAndReportsEmpty is the streaming half: a channel +// of reasoning-only chunks must be reported as empty so chainDrive degrades. +func TestPeekStreamHoldsReasoningAndReportsEmpty(t *testing.T) { + in := make(chan types.UnifiedChunk, 8) + for i := 0; i < 5; i++ { + in <- types.UnifiedChunk{ReasoningContent: "thinking…"} + } + in <- types.UnifiedChunk{Done: true, FinishReason: "length"} + close(in) + + _, peek := peekStream(in) + if peek() { + t.Error("a reasoning-only stream must be reported empty") + } +} + +// TestPeekStreamForwardsContentAfterReasoning: the common case still works, +// and the reasoning preamble is not forwarded ahead of the content (holding it +// back is what keeps the degrade path available). +func TestPeekStreamForwardsContentAfterReasoning(t *testing.T) { + in := make(chan types.UnifiedChunk, 8) + in <- types.UnifiedChunk{ReasoningContent: "let me think"} + in <- types.UnifiedChunk{Content: "Hi"} + in <- types.UnifiedChunk{Content: "!"} + in <- types.UnifiedChunk{Done: true, FinishReason: "stop"} + close(in) + + out, peek := peekStream(in) + if !peek() { + t.Fatal("a stream with content must be reported as having content") + } + var got []string + for ck := range out { + if ck.Content != "" { + got = append(got, ck.Content) + } + if strings.TrimSpace(ck.ReasoningContent) != "" { + t.Error("reasoning preamble must not be forwarded before content; " + + "that is what pins the client to a stream it cannot escape") + } + } + if strings.Join(got, "") != "Hi!" { + t.Errorf("forwarded content = %q, want %q", got, "Hi!") + } +} + +// TestPeekStreamDoesNotDeadlockOnToolCalls: a tool-call delta is content for +// this purpose and must unblock peek immediately. +func TestPeekStreamDoesNotDeadlockOnToolCalls(t *testing.T) { + raw, _ := json.Marshal([]types.ToolCall{{Name: "ls"}}) + in := make(chan types.UnifiedChunk, 4) + in <- types.UnifiedChunk{ToolCalls: raw} + in <- types.UnifiedChunk{Done: true, FinishReason: "tool_calls"} + close(in) + + _, peek := peekStream(in) + done := make(chan bool, 1) + go func() { done <- peek() }() + select { + case ok := <-done: + if !ok { + t.Error("a tool-call delta must count as content") + } + case <-time.After(2 * time.Second): + t.Fatal("peek blocked on a tool-call-only stream") + } + _ = context.Background() +} diff --git a/internal/scheduler/scheduler.go b/internal/scheduler/scheduler.go index 63bc7b5..3f467ce 100644 --- a/internal/scheduler/scheduler.go +++ b/internal/scheduler/scheduler.go @@ -11,6 +11,7 @@ import ( "fmt" "sort" "strings" + "sync" "sync/atomic" "time" @@ -245,6 +246,115 @@ func normalCount(cands []candidate) int { // // Normal candidates rotate by base; probe candidates form a fixed tail tried // only after every normal slot failed or was busy. +// emptyResultReason describes why a 200-with-no-content response counts as a +// slot failure for AUTO. +// +// WHY: reasoning models (claude-opus-*, codebuddy_glm-*, …) emit +// `reasoning_content` first and only then `content`. When the caller's +// max_tokens is small enough that the thinking phase consumes the whole budget, +// upstream returns 200 / finish_reason=length with 28 chunks of reasoning and +// ZERO content. runTier treated "err == nil" as success and handed that to the +// client, which then failed with "returned a completed response with no +// content" — a client-side error message for what is really a bad slot choice. +// +// So an empty result is a SLOT failure, not a request failure: the gateway +// degrades to the next slot and the user still gets an answer. Measured on +// production AUTO: 2 of 20 requests returned empty content, all of them +// claude-opus-4-8. +// +// A response carrying tool_calls or image data is NOT empty: an agent turn +// legitimately produces tool calls with no text. ReasoningContent does NOT +// rescue it either — see resultIsEmpty. +const emptyResultReason = "upstream returned no content (reasoning-only response, or the token budget was consumed before any text)" + +// resultIsEmpty reports whether a successful-but-useless response should be +// treated as a slot failure. +// +// Image data counts: an image-generation slot legitimately returns no text. +// A usage-only response is NOT empty either — the upstream answered, it just +// said nothing, and that is exactly the case worth degrading away from. +func resultIsEmpty(resp *types.UnifiedResponse) bool { + if resp == nil { + return false + } + // NOTE: ReasoningContent is deliberately NOT consulted. My first version + // excluded it ("the model was thinking, that is an answer"), and the test + // built from the real production capture failed immediately: the captured + // response is exactly reasoning_content-with-usage and zero text. The + // client asked for text and there is none; holding a request hostage to + // another model's thinking phase is strictly worse than degrading. + return strings.TrimSpace(resp.Content) == "" && + len(resp.ToolCalls) == 0 && + len(resp.ImageData) == 0 +} + +// emptyStreamReason is resultIsEmpty's streaming twin; see emptyResultReason +// for why an empty response is a slot failure rather than a request failure. +const emptyStreamReason = emptyResultReason + +// peekStream wraps a chunk channel so the caller learns whether the stream +// produced real content BEFORE the chunks are forwarded. +// +// Why this is necessary: reasoning models emit reasoning_content first. With a +// small max_tokens the whole budget is spent thinking, the stream ends with +// finish_reason=length and zero content. If the gateway forwarded those chunks +// as they arrived, the client would already have seen a 200 SSE stream and +// could not be given a different slot — its only recourse is the useless +// "returned a completed response with no content" error. Buffering until the +// first real content (or the end of the stream) keeps the degrade path +// available at the cost of holding back the first few chunks. +// +// What is NOT buffered: the wrapper starts forwarding as soon as a chunk with +// non-empty Content or ToolCalls arrives, and keeps forwarding everything from +// then on, so only the reasoning preamble is held. Reasoning-only responses +// are dropped in full and reported as empty, which lets chainDrive try the +// next slot. +func peekStream(in <-chan types.UnifiedChunk) (<-chan types.UnifiedChunk, func() bool) { + out := make(chan types.UnifiedChunk, 16) + var ( + mu sync.Mutex + sawText bool + done bool + ) + go func() { + defer close(out) + started := false + for ck := range in { + if !started { + // Hold back the reasoning / usage-only preamble. A tool-call + // delta counts as content: an agent turn legitimately emits + // tool_calls with no text. + if strings.TrimSpace(ck.Content) == "" && len(ck.ToolCalls) == 0 { + continue + } + started = true + mu.Lock() + sawText = true + mu.Unlock() + } + out <- ck + } + mu.Lock() + done = true + mu.Unlock() + }() + // peek blocks until the stream either produces content or ends, then + // reports whether any content was seen. Polling a 2ms tick rather than + // using a second channel keeps peekStream single-goroutine and leak-free. + peek := func() bool { + for { + mu.Lock() + seen, finished := sawText, done + mu.Unlock() + if seen || finished { + return seen + } + time.Sleep(2 * time.Millisecond) + } + } + return out, peek +} + func runTier(ctx context.Context, tn *TierNode, cands []candidate, base int64, req *types.ChatRequest, stream bool) tierResult { n := len(cands) norm := normalCount(cands) @@ -267,7 +377,18 @@ func runTier(ctx context.Context, tn *TierNode, cands []candidate, base int64, r if stream { chunks, err := sl.Prov.ChatStream(ctx, &r) if err == nil { - return tierResult{chunks: chunks, src: sl.Source, model: sl.Model} + guarded, peek := peekStream(chunks) + if peek() { + return tierResult{chunks: guarded, src: sl.Source, model: sl.Model} + } + // The stream finished with no content at all: a + // reasoning-only response. Drain and move on to the next + // slot instead of pinning the client to a useless stream. + hard = append(hard, TierError{ + Tier: tn.Tier, Source: sl.Source, Model: sl.Model, + Err: errors.New(emptyStreamReason), + }) + continue } if ctx.Err() != nil { return tierResult{} @@ -280,6 +401,17 @@ func runTier(ctx context.Context, tn *TierNode, cands []candidate, base int64, r } resp, err := sl.Prov.Chat(ctx, &r) if err == nil { + if resultIsEmpty(resp) { + // Soft failure: record it and try the next slot. Deliberately + // NOT a hard TierError — a hard error is reported to the client + // verbatim when the whole chain fails, and "this one model was + // unhelpful" is not the client's problem to debug. + hard = append(hard, TierError{ + Tier: tn.Tier, Source: sl.Source, Model: sl.Model, + Err: errors.New(emptyResultReason), + }) + continue + } return tierResult{resp: resp, src: sl.Source, model: sl.Model} } if ctx.Err() != nil { @@ -293,13 +425,75 @@ func runTier(ctx context.Context, tn *TierNode, cands []candidate, base int64, r return tierResult{hard: hard} } +// TraceKind classifies one step of an AUTO chain walk. +type TraceKind string + +const ( + // TraceTierSkip: the whole tier was skipped — every slot was cooling, + // quota-exhausted, or none was schedulable. Reason says which. + TraceTierSkip TraceKind = "tier_skip" + // TraceSlotFail: one slot failed hard (upstream error / bad adapter). The + // walk continues to the next slot or tier. + TraceSlotFail TraceKind = "slot_fail" + // TraceTierBusy: the tier was fully busy and the bounded wait expired. + TraceTierBusy TraceKind = "tier_busy" + // TraceSelected: this slot served the request. Exactly one per successful + // chain walk, and the last event emitted. + TraceSelected TraceKind = "selected" +) + +// TraceEvent is one observable step of an AUTO chain walk. +// +// WHY THIS EXISTS: chainDrive's return value is (resp, src, model, err), so a +// caller learns only which slot finally served the request. Everything the +// scheduler decided on the way there — which tiers it skipped and WHY, which +// slots hard-failed, whether a tier was merely busy — was computed and then +// discarded. That is invisible to operators and to plugins: "tier 1 was cooling +// so we degraded to tier 3" looked exactly like "tier 1 served it". +// +// The walk already accumulates this in ChainErr, but ONLY on total failure, and +// ChainErr is an error return, not a record. Emitting a trace as it happens +// covers the far more common case: a request that SUCCEEDED after degrading. +// +// Design constraints: +// - scheduler stays dependency-free and independently testable. A TraceEvent +// is a plain struct in this package and the sink is a func parameter, so no +// import is added and no test has to change to observe a walk. +// - The sink is optional (nil = emit nothing). The overhead on the hot path +// is one nil check per event. +// - Events are OBSERVATION ONLY. Nothing in the scheduler branches on them, +// and the gateway does not feed them back into routing, cooldown or quota — +// see docs/plugins.md for why accounting and enforcement are kept apart. +type TraceEvent struct { + Kind TraceKind + Tier int + Source string + Model string + Reason string // human-readable, for TraceTierSkip / TraceSlotFail + Err string // the underlying error text, for TraceSlotFail + // Attempt counts the 1-based slot attempt within the whole walk. + Attempt int +} + +// TraceSink receives chain-walk events. It must not block: it is called from the +// request path, and a slow sink slows the request. +type TraceSink func(TraceEvent) + // chainDrive runs a request down the chain (plan 2.3): tiers ascending (tier // 1, the highest priority, first), per-tier round-robin starting at the tier // cursor, same-tier runs ordered by preference (negative prefs sink but stay // reachable). Quota-exhausted and cooling slots are filtered up front; a // fully busy tier is polled for a bounded time before falling through. // Failures are summarized in *ChainErr for the caller to map to HTTP 503. -func (s *Scheduler) chainDrive(ctx context.Context, chain *Chain, req *types.ChatRequest, exhausted func(*Slot) bool, stream bool) (*types.UnifiedResponse, <-chan types.UnifiedChunk, string, string, error) { +// +// trace may be nil; when set it receives one event per observable step. +func (s *Scheduler) chainDrive(ctx context.Context, chain *Chain, req *types.ChatRequest, exhausted func(*Slot) bool, stream bool, trace TraceSink) (*types.UnifiedResponse, <-chan types.UnifiedChunk, string, string, error) { + emit := func(ev TraceEvent) { + if trace != nil { + trace(ev) + } + } + attempt := 0 if chain == nil || len(chain.Tiers) == 0 { return nil, nil, "", "", fmt.Errorf("no auto slot configured") } @@ -309,7 +503,9 @@ func (s *Scheduler) chainDrive(ctx context.Context, chain *Chain, req *types.Cha // dropped unless they qualify as half-cooldown probes (appended last). cands := collectCands(tn.Slots, exhausted) if len(cands) == 0 { - ce.Skipped = append(ce.Skipped, fmt.Sprintf("tier %d: no schedulable slot (cooling or quota exhausted)", tn.Tier)) + reason := "no schedulable slot (cooling or quota exhausted)" + ce.Skipped = append(ce.Skipped, fmt.Sprintf("tier %d: %s", tn.Tier, reason)) + emit(TraceEvent{Kind: TraceTierSkip, Tier: tn.Tier, Reason: reason}) continue } // No Pref sort: load balancing is done by round-robin cursor. @@ -318,6 +514,8 @@ func (s *Scheduler) chainDrive(ctx context.Context, chain *Chain, req *types.Cha base := tn.NextStart() res := runTier(ctx, tn, cands, base, req, stream) if res.resp != nil || res.chunks != nil { + attempt++ + emit(TraceEvent{Kind: TraceSelected, Tier: tn.Tier, Source: res.src, Model: res.model, Attempt: attempt}) releaseProbes(cands) return res.resp, res.chunks, res.src, res.model, nil } @@ -327,6 +525,13 @@ func (s *Scheduler) chainDrive(ctx context.Context, chain *Chain, req *types.Cha } if len(res.hard) > 0 { ce.Tiers = append(ce.Tiers, res.hard...) + for _, h := range res.hard { + attempt++ + emit(TraceEvent{ + Kind: TraceSlotFail, Tier: tn.Tier, Source: h.Source, Model: h.Model, + Err: types.OneLine(h.Err.Error(), 200), Attempt: attempt, + }) + } releaseProbes(cands) continue // hard failures: fall through to the next tier, no waiting } @@ -334,10 +539,13 @@ func (s *Scheduler) chainDrive(ctx context.Context, chain *Chain, req *types.Cha if err := s.pollBusyTier(ctx, tn, cands, base, req, stream, &ce); err != nil { releaseProbes(cands) if r, ok := err.(*tierSuccess); ok { + attempt++ + emit(TraceEvent{Kind: TraceSelected, Tier: tn.Tier, Source: r.res.src, Model: r.res.model, Attempt: attempt}) return r.res.resp, r.res.chunks, r.res.src, r.res.model, nil } return nil, nil, "", "", err } + emit(TraceEvent{Kind: TraceTierBusy, Tier: tn.Tier, Reason: fmt.Sprintf("no free slot within %v", busyWait)}) releaseProbes(cands) } if len(ce.Tiers) == 0 && len(ce.Skipped) == 0 { @@ -401,16 +609,16 @@ func (s *Scheduler) pollBusyTier(ctx context.Context, tn *TierNode, cands []cand // non-nil, decides slot token-quota exhaustion. Returns the response, the // serving source and the exact model id used; on total failure a *ChainErr // summarizing every tier. -func (s *Scheduler) ChainChat(ctx context.Context, chain *Chain, req *types.ChatRequest, exhausted func(*Slot) bool) (*types.UnifiedResponse, string, string, error) { - resp, _, src, model, err := s.chainDrive(ctx, chain, req, exhausted, false) +func (s *Scheduler) ChainChat(ctx context.Context, chain *Chain, req *types.ChatRequest, exhausted func(*Slot) bool, trace TraceSink) (*types.UnifiedResponse, string, string, error) { + resp, _, src, model, err := s.chainDrive(ctx, chain, req, exhausted, false, trace) return resp, src, model, err } // ChainChatStream runs a streaming AUTO request down the chain. A slot is // abandoned only on connect failures / busy (before its first chunk); after a // stream starts it is pinned. Same return contract as ChainChat. -func (s *Scheduler) ChainChatStream(ctx context.Context, chain *Chain, req *types.ChatRequest, exhausted func(*Slot) bool) (<-chan types.UnifiedChunk, string, string, error) { - _, chunks, src, model, err := s.chainDrive(ctx, chain, req, exhausted, true) +func (s *Scheduler) ChainChatStream(ctx context.Context, chain *Chain, req *types.ChatRequest, exhausted func(*Slot) bool, trace TraceSink) (<-chan types.UnifiedChunk, string, string, error) { + _, chunks, src, model, err := s.chainDrive(ctx, chain, req, exhausted, true, trace) return chunks, src, model, err } diff --git a/internal/scheduler/scheduler_test.go b/internal/scheduler/scheduler_test.go index c7c0a6f..96811fc 100644 --- a/internal/scheduler/scheduler_test.go +++ b/internal/scheduler/scheduler_test.go @@ -176,7 +176,7 @@ func TestChainRoundRobin(t *testing.T) { s := New(0) var got []string for i := 0; i < 4; i++ { - _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil) + _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) if err != nil { t.Fatalf("iter %d: %v", i, err) } @@ -200,7 +200,7 @@ func TestChainPreferenceSinksButStaysReachable(t *testing.T) { {Tier: 0, Model: "g", Source: "good"}, }, bySource(neg, good)) s := New(0) - resp, src, model, err := s.ChainChat(context.Background(), ch, chatReq(), nil) + resp, src, model, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) if err != nil { t.Fatalf("chain: %v", err) } @@ -213,7 +213,7 @@ func TestChainPreferenceSinksButStaysReachable(t *testing.T) { t.Fatal("neg was tried first and succeeded; good must not be attempted") } // Second request: cursor advances. neg wins again (good hard-fails). - resp2, src2, _, err2 := s.ChainChat(context.Background(), ch, chatReq(), nil) + resp2, src2, _, err2 := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) if err2 != nil { t.Fatalf("second chain: %v", err2) } @@ -230,7 +230,7 @@ func TestChainBusySkipsWithoutPenalty(t *testing.T) { {Tier: 0, Model: "b", Source: "s2"}, }, bySource(a, b)) s := New(0) - _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil) + _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) if err != nil { t.Fatalf("chain: %v", err) } @@ -256,7 +256,7 @@ func TestChainAllBusyBoundedWaitThenNextTier(t *testing.T) { }, bySource(a, b, c)) s := New(0) t0 := time.Now() - _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil) + _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) el := time.Since(t0) if err != nil { t.Fatalf("chain: %v", err) @@ -277,7 +277,7 @@ func TestChainQuotaExhausted(t *testing.T) { }, bySource(a, b)) s := New(0) exhausted := func(sl *Slot) bool { return sl.Source == "s1" && sl.Quota > 0 } - _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), exhausted) + _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), exhausted, nil) if err != nil { t.Fatalf("chain: %v", err) } @@ -300,7 +300,7 @@ func TestChainErrSummary(t *testing.T) { {Tier: 1, Model: "c", Source: "s3"}, }, bySource(a, b, c)) s := New(0) - _, _, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil) + _, _, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) var ce *ChainErr if !errors.As(err, &ce) { t.Fatalf("err = %v, want *ChainErr", err) @@ -325,7 +325,7 @@ func TestChainStreamFallsBackBeforeFirstChunk(t *testing.T) { {Tier: 0, Model: "b", Source: "s2"}, }, bySource(a, b)) s := New(0) - chunks, src, model, err := s.ChainChatStream(context.Background(), ch, chatReq(), nil) + chunks, src, model, err := s.ChainChatStream(context.Background(), ch, chatReq(), nil, nil) if err != nil { t.Fatalf("chain stream: %v", err) } @@ -365,7 +365,7 @@ func TestChainProbeIsLastResort(t *testing.T) { }, bySource(healthy, cooling)) s := New(0) for i := 0; i < 3; i++ { - _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil) + _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) if err != nil { t.Fatalf("iter %d: %v", i, err) } @@ -390,7 +390,7 @@ func TestChainProbeServesWhenNothingElseCan(t *testing.T) { cooling.probeable.Store(true) ch := BuildChain([]Rule{{Tier: 0, Model: "c", Source: "cooling"}}, bySource(cooling)) s := New(0) - resp, src, model, err := s.ChainChat(context.Background(), ch, chatReq(), nil) + resp, src, model, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) if err != nil { t.Fatalf("probe must serve the request: %v", err) } @@ -417,14 +417,14 @@ func TestChainProbePermitReleasedOnFailure(t *testing.T) { cooling.fail.Store(true) ch := BuildChain([]Rule{{Tier: 0, Model: "c", Source: "cooling"}}, bySource(cooling)) s := New(0) - if _, _, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil); err == nil { + if _, _, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil); err == nil { t.Fatal("expected the failing probe to surface an error") } if cooling.probeClaims.Load() != 1 || cooling.probeDones.Load() != 1 { t.Fatalf("permit accounting: claims=%d dones=%d, want 1/1", cooling.probeClaims.Load(), cooling.probeDones.Load()) } // permit is free again for the next attempt - if _, _, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil); err == nil { + if _, _, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil); err == nil { t.Fatal("expected the second probe to fail too") } if cooling.probeClaims.Load() != 2 { @@ -443,7 +443,7 @@ func TestChainProbeDoesNotBlockTierFallthrough(t *testing.T) { {Tier: 2, Model: "b", Source: "backup"}, }, bySource(cold, backup)) s := New(0) - _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil) + _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) if err != nil || src != "backup" { t.Fatalf("want fallthrough to backup, got src=%q err=%v", src, err) } @@ -482,7 +482,7 @@ func TestChainProbeRoundRobinUnaffected(t *testing.T) { s := New(0) var got []string for i := 0; i < 4; i++ { - _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil) + _, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil, nil) if err != nil { t.Fatalf("iter %d: %v", i, err) } diff --git a/internal/scheduler/trace_test.go b/internal/scheduler/trace_test.go new file mode 100644 index 0000000..2873e42 --- /dev/null +++ b/internal/scheduler/trace_test.go @@ -0,0 +1,191 @@ +package scheduler + +import ( + "context" + "errors" + "strings" + "testing" +) + +// The chain trace is the only way a caller learns that a request was DEGRADED +// — served by a lower tier than the one that should have taken it. chainDrive's +// return value carries only the winner, so without these events "tier 1 was +// cooling and we dropped to tier 2" is indistinguishable from "tier 1 served +// it", which is the exact question a priority chain exists to answer. + +// recorder collects trace events for assertions. +type recorder struct{ events []TraceEvent } + +func (r *recorder) sink(ev TraceEvent) { r.events = append(r.events, ev) } + +func (r *recorder) kinds() []TraceKind { + out := make([]TraceKind, 0, len(r.events)) + for _, e := range r.events { + out = append(out, e.Kind) + } + return out +} + +func (r *recorder) find(k TraceKind) *TraceEvent { + for i := range r.events { + if r.events[i].Kind == k { + return &r.events[i] + } + } + return nil +} + +// TestTraceSelectedOnlyOnHappyPath: a clean walk emits exactly one event. +func TestTraceSelectedOnlyOnHappyPath(t *testing.T) { + p1 := fakeProv("p1", "m1") + ch := BuildChain([]Rule{{Model: "m1", Source: "p1", Tier: 1}}, + func(model, source string) Provider { return p1 }) + var rec recorder + _, _, _, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, rec.sink) + if err != nil { + t.Fatalf("chain: %v", err) + } + if got := rec.kinds(); len(got) != 1 || got[0] != TraceSelected { + t.Errorf("events = %v, want a single selected", got) + } + e := rec.find(TraceSelected) + if e.Tier != 1 || e.Source != "p1" || e.Model != "m1" { + t.Errorf("selected event = %+v, want tier 1 / p1 / m1", e) + } + if e.Attempt != 1 { + t.Errorf("Attempt = %d, want 1", e.Attempt) + } +} + +// TestTraceRecordsTierSkipAndDegradation is the core case: tier 1 is +// unschedulable, tier 2 answers. The trace must show the skip AND the eventual +// selection, so a consumer can see the request was served one tier down. +func TestTraceRecordsTierSkipAndDegradation(t *testing.T) { + // p1 is unavailable (not probeable), so tier 1 yields no candidates. + p1 := fakeProv("p1", "m1") + p1.available.Store(false) + p2 := fakeProv("p2", "m2") + ch := BuildChain([]Rule{ + {Model: "m1", Source: "p1", Tier: 1}, + {Model: "m2", Source: "p2", Tier: 2}, + }, bySource(p1, p2)) + var rec recorder + _, src, model, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, rec.sink) + if err != nil { + t.Fatalf("chain: %v", err) + } + if src != "p2" || model != "m2" { + t.Fatalf("served by %s/%s, want p2/m2", src, model) + } + skip := rec.find(TraceTierSkip) + if skip == nil { + t.Fatalf("no tier_skip event; events = %v", rec.kinds()) + } + if skip.Tier != 1 { + t.Errorf("skip tier = %d, want 1", skip.Tier) + } + if !strings.Contains(skip.Reason, "cooling") { + t.Errorf("skip reason = %q, want it to mention cooling", skip.Reason) + } + sel := rec.find(TraceSelected) + if sel == nil || sel.Tier != 2 { + t.Errorf("selected = %+v, want tier 2", sel) + } + // The order matters: the skip must be observable BEFORE the selection. + if rec.events[0].Kind != TraceTierSkip || rec.events[len(rec.events)-1].Kind != TraceSelected { + t.Errorf("event order = %v, want skip first and selected last", rec.kinds()) + } +} + +// TestTraceRecordsHardSlotFailures: a slot that returns an upstream error is a +// different event from a skip — the request tried it and it failed. Losing that +// distinction makes a flaky upstream look like an idle one. +func TestTraceRecordsHardSlotFailures(t *testing.T) { + p1 := fakeProv("p1", "m1") + p1.fail.Store(true) // Chat returns "upstream error" + p2 := fakeProv("p2", "m2") + ch := BuildChain([]Rule{ + {Model: "m1", Source: "p1", Tier: 1}, + {Model: "m2", Source: "p2", Tier: 2}, + }, bySource(p1, p2)) + var rec recorder + _, src, _, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, rec.sink) + if err != nil { + t.Fatalf("chain: %v", err) + } + if src != "p2" { + t.Fatalf("served by %s, want p2", src) + } + fail := rec.find(TraceSlotFail) + if fail == nil { + t.Fatalf("no slot_fail event; events = %v", rec.kinds()) + } + if fail.Tier != 1 || fail.Source != "p1" || fail.Model != "m1" { + t.Errorf("slot_fail = %+v, want tier 1 / p1 / m1", fail) + } + if !strings.Contains(fail.Err, "upstream error") { + t.Errorf("slot_fail error = %q, want the upstream text", fail.Err) + } + // A hard failure must NOT be reported as a skip. + if rec.find(TraceTierSkip) != nil { + t.Error("a hard failure was also reported as a tier_skip") + } +} + +// TestTraceNilSinkIsSafe: the gateway passes nil when no plugin is loaded, so +// every emit path must tolerate it. This is the "plugins are optional" property +// on the scheduler side. +func TestTraceNilSinkIsSafe(t *testing.T) { + p1 := fakeProv("p1", "m1") + p1.fail.Store(true) + p2 := fakeProv("p2", "m2") + ch := BuildChain([]Rule{ + {Model: "m1", Source: "p1", Tier: 1}, + {Model: "m2", Source: "p2", Tier: 2}, + }, bySource(p1, p2)) + if _, _, _, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, nil); err != nil { + t.Fatalf("a nil trace sink broke the walk: %v", err) + } +} + +// TestTraceOnTotalFailure: when every tier fails, the walk still emits its +// per-step events AND returns the ChainErr. The trace is additive — it must not +// replace or disturb the error contract callers depend on for the 503. +func TestTraceOnTotalFailure(t *testing.T) { + p1 := fakeProv("p1", "m1") + p1.fail.Store(true) + p2 := fakeProv("p2", "m2") + p2.fail.Store(true) + ch := BuildChain([]Rule{ + {Model: "m1", Source: "p1", Tier: 1}, + {Model: "m2", Source: "p2", Tier: 2}, + }, bySource(p1, p2)) + var rec recorder + _, _, _, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, rec.sink) + var ce *ChainErr + if !errors.As(err, &ce) { + t.Fatalf("err = %v, want a *ChainErr so the gateway can answer 503", err) + } + if len(ce.Tiers) != 2 { + t.Errorf("ChainErr.Tiers = %d, want 2 (the error contract must be unchanged)", len(ce.Tiers)) + } + if n := len(rec.kinds()); n != 2 { + t.Errorf("events = %v, want two slot_fail and no selection", rec.kinds()) + } + if rec.find(TraceSelected) != nil { + t.Error("a selected event was emitted for a walk that served nothing") + } +} + +// TestTraceSkipsEmptyChain: no chain configured must not emit anything; the +// gateway answers 503 before scheduling in that case anyway. +func TestTraceSkipsEmptyChain(t *testing.T) { + var rec recorder + _, _, _, err := New(3).ChainChat(context.Background(), &Chain{}, chatReq(), nil, rec.sink) + if err == nil { + t.Fatal("expected an error for an empty chain") + } + if len(rec.events) != 0 { + t.Errorf("events = %v, want none", rec.kinds()) + } +} diff --git a/packaging/config.example.yaml b/packaging/config.example.yaml index a06d732..f1ef9c6 100644 --- a/packaging/config.example.yaml +++ b/packaging/config.example.yaml @@ -13,13 +13,26 @@ listen: 127.0.0.1:8080 # 客户端访问本网关所需的 API Key(Bearer)。留空数组 = 不鉴权(仅内网)。 gateway_keys: [] -# 默认模型选择:具体模型 id 或 AUTO(按各源模型的 priority 自动选最高可用源) +# 默认模型选择:具体模型 id 或 AUTO。 +# AUTO = 走 WebUI「优先级页」保存的调度链(存在本文件的 `auto:` 字段)。 default_model: AUTO # Lua 适配器目录(默认 adapters/,首次启动自动写入内置适配器) adapter_dir: adapters -# 运行时持久化文件(WebUI 新增/编辑的源会写入此文件,重启后仍生效) +# 插件目录(可选)。设置后: +# - 首次启动会把随核心发布的示例插件(billing:按源/模型/密钥计费 + 仪表盘) +# 写入本目录,它会在请求流水线上挂钩子,并向 WebUI 注入一个「Billing」页面 +# 和状态页上的一块总开销组件; +# - 之后 WebUI「插件」页与 Electron 壳的设置面板可安装/禁用/删除/编辑。 +# 目录一旦存在即以目录为准:删除或改写内置插件都是真实生效的操作。 +# 留空 = 插件功能完全关闭(不影响网关其它功能)。 +# plugin_dir: /etc/llmsproxy/plugins + +# 运行时文件:存放 WebUI 管理的源模板、已删除标记、预置模板名单。 +# +# 注意:**AUTO 调度链、网关密钥、上游源都存在本 config.yaml 里**, +# 不在这个文件。runtime.json 只管模板与删除标记。 runtime_file: runtime.json # 全局并发上限(0 = 不限) diff --git a/packaging/llmsproxy.service b/packaging/llmsproxy.service index f64d507..90b846d 100644 --- a/packaging/llmsproxy.service +++ b/packaging/llmsproxy.service @@ -10,15 +10,53 @@ Type=simple # OS thread that touches malloc reserved its own ~1 MB arena that is never # returned. Measured: 8-12 arenas -> 0. # GOGC=50 halves the Go heap growth target. On its own it does NOT help (the -# saved heap is immediately eaten by more glibc arenas); combined with +# saved heap is immediately eaten by extra glibc arenas); combined with # MALLOC_ARENA_MAX it cut settled RSS by ~19%. This gateway is I/O bound, so # the extra GC cycles are free. Environment=GOGC=50 Environment=MALLOC_ARENA_MAX=2 -ExecStart=/usr/bin/llmsproxy -config /etc/llmsproxy/config.yaml +ExecStart=/usr/local/bin/llmsproxy -config /etc/llmsproxy/config.yaml WorkingDirectory=/etc/llmsproxy Restart=always RestartSec=5 +# ---- 加固(2026-10-01 逐条实测后加入,不是照抄文档)---- +# +# 为什么仍然以 root 运行:master.key 是 0600 root。加 User=llmsproxy 实测直接 +# 起不来,而且失败方式很隐蔽—— +# [config] secrets disabled: open /etc/llmsproxy/master.key: permission denied +# 只是**一行日志**,服务会带着"敏感值将以明文落盘"继续跑起来。 +# 也就是说降权在当前文件权限下不是加固,而是把密钥降级。要降权必须先把 +# master.key 交给服务用户并统一 /etc/llmsproxy 的属主,那是一次独立的、有回滚 +# 需求的变更,不该和加固混在一起。 +# +# 下面每一条都在一个独立探针单元(临时端口 + 独立 runtime_file/adapter_dir) +# 上真实验证过:鉴权 401/200 正常、发一次真实 /v1/chat/completions 走通(证明 +# LuaJIT 适配器路径没被 seccomp 打断)、审计文件可写可轮转、连续重启 3 次与 +# kill -9 后行为符合预期。systemd 对非法指令值不报错只"忽略",逐条实测是唯一 +# 可靠做法。 +NoNewPrivileges=yes +# 读路径全部落在 /etc/llmsproxy;写路径经核对只有 config.yaml / runtime.json / +# audit.jsonl / adapters/*.lua / master.key,全在该目录下(internal/{config,gateway, +# core,lua} 里的 WriteFile|Rename|Remove 调用点)。 +ProtectSystem=strict +ReadWritePaths=/etc/llmsproxy +ProtectHome=yes +PrivateTmp=yes +ProtectKernelTunables=yes +ProtectKernelModules=yes +ProtectControlGroups=yes +RestrictSUIDSGID=yes +RestrictRealtime=yes +LockPersonality=yes +RestrictAddressFamilies=AF_INET AF_INET6 AF_UNIX +# CapabilityBoundingSet 置空:本服务不需要任何 capability(不建 netns、不改 +# 资源限制、不 chown)。留空即"一个都不给",比列一份允许清单更难写错。 +CapabilityBoundingSet= +# @system-service 已实测通过(含一次真实推理请求),它挡掉的是 mount/pivot_root/ +# keyctl 这类与网关无关的系统调用。 +SystemCallFilter=@system-service +SystemCallArchitectures=native + [Install] WantedBy=multi-user.target \ No newline at end of file