Files
ModelRouter/internal/lua/persist_test.go
JianFeeeee ce2032435a fix(plugin): 插件 state 持久化 —— 重启不再丢账
## 问题(压测实测)
插件 state 活在 Lua VM 里,进程一死就没了。实测线上量级:
  重启前 {"requests":218241,"prompt_tokens":26188920,...}
  重启后 {"requests":0,"prompt_tokens":0,...}
对计费插件来说这不是舍入误差,是功能本身没生效 —— 它存在的意义就是那个
不断累加的数字,而一次 systemctl restart 就能把它抹掉。

## 设计
- prices(配置)与 state(累计历史)**分开存**在同一个文件里但不同字段。
  SetState 在内存里已经这么分,磁盘必须同意:合并会让改价看起来像清零,
  或让恢复历史时顺带复活过期价格。
- 原子写(临时文件 + rename):写一半崩掉时上一份仍可读,而不是留下一个
  解析失败的半截 JSON —— 那等于这次也丢。
- 损坏文件只警告不阻断启动。转发不能依赖插件的账本活着。
- 防抖后台刷:钩子路径只标记,真正的写在一个 goroutine 里合并进行。
  计费插件每请求都改 state,同步写会把一次 JSON 编码 + 文件写放到热路径上
  (实测钩子本身已经 14.6µs,写会盖过它)。
- Core.Close 必须先刷插件再停 VM:flush 要读 Lua 表,vm.Stop() 之后读的是
  已释放的内存。

## ★ 实现中踩的四个坑(都由测试或崩溃直接暴露,不是推测)
1. **后台 goroutine 碰 Lua = use-after-free**。最初让 flush 线程去读 Lua 状态,
   vm.Stop() 后那是已释放内存 —— 表现为 golua 里的 SIGSEGV,不是干净报错。
   改成:钩子路径(VM 必然存活、已持 p.mu)取快照,后台只写文件。
2. **自死锁**:markDirtyLocked 被 invoke 调用,而 invoke 全程持 p.mu,
   再 Lock 一次就是死锁。lua 包测试直接挂到超时。
3. **luaToJSON 独占整个栈**(每条路径结尾都 SetTop(0))。连续调两次读两个
   字段时第二次访问的是不存在的槽位 —— 这个绑定不 panic,直接 SIGABRT。
   改为每次重建栈。中间还因为提前 return 没 Pop 而让栈逐次错位。
4. **快照顺序**:先快照后读返回值,会把钩子的返回值清掉,于是每个"有意见"的
   插件静默变成"没意见",而文档承诺的"返回 table 合并进 payload"就废了,
   且没有任何报错。

## 判据(7 项,全部变异验证过)
重启后总计保留 / prices 与 state 分离 / 纯 prices 更新也持久化 /
钩子返回值不被快照吃掉 / 损坏文件降级不阻断 / 500 次变更合并成个位数次写 /
Close 刷出尾部。

变异结果:
  关掉 mark          → TestStateSurvivesRestart + TestPricesAndStateAreSeparate 红
  Close 不等 flush   → TestCloseFlushesTail 红
  prices-only 不写盘 → TestPricesOnlyUpdatePersists 红
  还原快照顺序       → TestHookReturnValueSurvivesSnapshot 红
★ 第一次跑「关掉 mark」时判据没报错,原因是我的变异脚本写出未使用变量导致
  编译失败 —— go test 根本没跑测试,我却读成了"通过"。换成 _, _ = 后如期变红。

## 端到端
隔离实例发 12 次请求 → systemctl restart → requests 仍为 12,token 数不变。
371 个测试全绿。
2026-10-02 10:17:07 +08:00

284 lines
8.7 KiB
Go

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"])
}
}