mirror of
https://gitcode.com/JianFeeeee/ModelRouter.git
synced 2026-10-03 23:54:06 +00:00
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 个测试全绿。
This commit is contained in:
@ -826,6 +826,13 @@ func normalizeSource(s *config.Source) error {
|
||||
|
||||
// Close releases resources.
|
||||
func (c *Core) Close() {
|
||||
// Plugin state must be flushed BEFORE the VM stops. The saver's final write
|
||||
// reads each plugin's Lua tables; once vm.Stop() has closed those states the
|
||||
// read finds nothing and the last interval of accumulation is lost — which
|
||||
// is the exact failure this persistence was added to prevent.
|
||||
if c.plugins != nil {
|
||||
c.plugins.Close()
|
||||
}
|
||||
if c.vm != nil {
|
||||
c.vm.Stop()
|
||||
}
|
||||
|
||||
283
internal/lua/persist_test.go
Normal file
283
internal/lua/persist_test.go
Normal file
@ -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"])
|
||||
}
|
||||
}
|
||||
@ -40,6 +40,7 @@ import (
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
golua "github.com/aarzilli/golua/lua"
|
||||
)
|
||||
@ -194,6 +195,12 @@ type Plugin struct {
|
||||
// which is the real cost — documented as a rule in docs/plugins.md.
|
||||
mu sync.Mutex
|
||||
state *worker
|
||||
|
||||
// persisted marks this plugin's state as having been restored from disk.
|
||||
// It gates the first save: without it, a freshly loaded plugin whose state
|
||||
// is still the compiled-in default would immediately overwrite the file it
|
||||
// was supposed to inherit from.
|
||||
persisted bool
|
||||
}
|
||||
|
||||
// Plugins is the loaded plugin set, owned by the VM.
|
||||
@ -209,6 +216,156 @@ type Plugins struct {
|
||||
// hookErr records per-stage plugin failures so a silently broken plugin is
|
||||
// visible in /api/status rather than merely missing.
|
||||
hookErr *hookErrors
|
||||
// logf, when set, receives persistence warnings. It is a field rather than a
|
||||
// direct log call so the plugin package stays free of a logging dependency
|
||||
// and tests can capture the warnings.
|
||||
logf func(format string, args ...interface{})
|
||||
// stateDirDisabled turns persistence off. Used by tests that assert state
|
||||
// starts empty, and by an embedder that has no writable plugin dir.
|
||||
stateDirDisabled bool
|
||||
// save coalesces state writes into one background goroutine.
|
||||
//
|
||||
// A hook must NOT write synchronously: the billing plugin mutates state on
|
||||
// every single request, and serializing the whole state per request would put
|
||||
// a file write (plus a full JSON encode) on the hot path — measured at
|
||||
// 14.6us per hook call already, a write would dominate it. Instead hooks mark
|
||||
// the plugin dirty and this saver flushes, so N requests between two flushes
|
||||
// cost one write.
|
||||
save *stateSaver
|
||||
}
|
||||
|
||||
// stateSaver coalesces state writes: one goroutine, a minimum interval between
|
||||
// flushes, and a dirty set. Under load the number of writes is bounded by the
|
||||
// timer rather than by the request rate.
|
||||
type stateSaver struct {
|
||||
mu sync.Mutex
|
||||
ps *Plugins
|
||||
dirty map[string]pendingState
|
||||
wake chan struct{}
|
||||
stop chan struct{}
|
||||
done chan struct{}
|
||||
// wakeDelay is how long a burst suppresses the next tick. A separate field
|
||||
// (not the constant 2s) so tests can shorten it: Close() waits for the saver
|
||||
// goroutine, so a test with the production interval pays that interval on
|
||||
// every teardown.
|
||||
wakeDelay time.Duration
|
||||
once sync.Once
|
||||
// interval is the minimum gap between flushes.
|
||||
interval time.Duration
|
||||
// writes counts completed file writes; tests read it to assert that N
|
||||
// mutations did NOT become N writes.
|
||||
writes atomic.Int64
|
||||
}
|
||||
|
||||
// pendingState is a plugin's state ALREADY CONVERTED to Go values.
|
||||
//
|
||||
// The conversion happens on the caller's goroutine, inside p.mu, while the Lua
|
||||
// VM is guaranteed alive. Doing it in the flush goroutine instead looked
|
||||
// harmless and was a use-after-free: vm.Stop() frees the Lua states, and the
|
||||
// saver then walked them — measured as a SIGSEGV inside golua, not as a clean
|
||||
// error. A background goroutine must never touch the VM.
|
||||
type pendingState struct {
|
||||
name string
|
||||
state interface{}
|
||||
prices interface{}
|
||||
}
|
||||
|
||||
// setWakeDelayForTest shortens the debounce so tests do not pay it on teardown.
|
||||
func (s *stateSaver) setWakeDelayForTest(d time.Duration) { s.wakeDelay = d }
|
||||
|
||||
func newStateSaver(ps *Plugins) *stateSaver {
|
||||
s := &stateSaver{
|
||||
ps: ps,
|
||||
dirty: map[string]pendingState{},
|
||||
wake: make(chan struct{}, 1),
|
||||
stop: make(chan struct{}),
|
||||
done: make(chan struct{}),
|
||||
interval: 2 * time.Second,
|
||||
wakeDelay: 2 * time.Second,
|
||||
}
|
||||
go s.loop()
|
||||
return s
|
||||
}
|
||||
|
||||
// mark records a snapshot as needing a write. Never blocks: the channel is
|
||||
// buffered, and a full buffer means a flush is already pending, which is
|
||||
// exactly the state we want. Only the LATEST snapshot per plugin is kept, so a
|
||||
// burst of N mutations collapses to one write.
|
||||
func (s *stateSaver) mark(p pendingState) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
s.dirty[p.name] = p
|
||||
s.mu.Unlock()
|
||||
select {
|
||||
case s.wake <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stateSaver) loop() {
|
||||
defer close(s.done)
|
||||
t := time.NewTicker(s.interval)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-s.stop:
|
||||
// Final flush so a clean shutdown does not lose the tail — losing
|
||||
// the last interval of spend is the same bug in a smaller window.
|
||||
s.flush()
|
||||
return
|
||||
case <-t.C:
|
||||
s.flush()
|
||||
case <-s.wake:
|
||||
// Debounce a burst of marks into one write.
|
||||
time.Sleep(s.wakeDelay)
|
||||
s.flush()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// flush writes every dirty plugin's state.
|
||||
func (s *stateSaver) flush() {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
if len(s.dirty) == 0 {
|
||||
s.mu.Unlock()
|
||||
return
|
||||
}
|
||||
pending := make([]pendingState, 0, len(s.dirty))
|
||||
for _, p := range s.dirty {
|
||||
pending = append(pending, p)
|
||||
}
|
||||
s.dirty = map[string]pendingState{}
|
||||
s.mu.Unlock()
|
||||
|
||||
// File I/O only. The Lua states are not touched here — see pendingState.
|
||||
for _, p := range pending {
|
||||
writeJSONAtomic(s.ps.stateFile(p.name), map[string]interface{}{
|
||||
"version": 1,
|
||||
"state": p.state,
|
||||
"prices": p.prices,
|
||||
})
|
||||
s.writes.Add(1)
|
||||
}
|
||||
}
|
||||
|
||||
// Close stops the saver AFTER its final flush completes.
|
||||
//
|
||||
// It must WAIT for that flush, not just signal it. Signalling and returning
|
||||
// leaves the write to a goroutine that the caller is about to tear down
|
||||
// (vm.Stop() frees the plugin states; the test process is exiting), so the
|
||||
// final state is silently lost — which is the very defect persistence was added
|
||||
// to fix. Close is on the shutdown path, where a few milliseconds is free.
|
||||
func (s *stateSaver) Close() {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.once.Do(func() { close(s.stop) })
|
||||
<-s.done
|
||||
}
|
||||
|
||||
type hookCall struct {
|
||||
@ -259,12 +416,30 @@ const pluginGlobal = "__llmsproxy_plugin"
|
||||
// NewPlugins creates the plugin registry for a VM. dir is the plugin directory;
|
||||
// a missing directory is not an error (plugins are optional).
|
||||
func NewPlugins(vm *VM, dir string) *Plugins {
|
||||
return &Plugins{
|
||||
ps := &Plugins{
|
||||
vm: vm,
|
||||
stageFuncs: map[Stage][]hookCall{},
|
||||
hookErr: newHookErrors(),
|
||||
dir: dir,
|
||||
}
|
||||
// The saver is started even with no dir: a gateway can be given a plugin dir
|
||||
// later, and a saver that only exists when dir != "" would silently never
|
||||
// flush. markDirtyLocked and flush both no-op when the dir is empty.
|
||||
ps.save = newStateSaver(ps)
|
||||
return ps
|
||||
}
|
||||
|
||||
// Close stops the background state saver, flushing once more first.
|
||||
//
|
||||
// The gateway MUST call this on shutdown. Without it the last flush interval's
|
||||
// accumulation is lost, which is the same "restart loses the books" defect this
|
||||
// persistence exists to fix, just scoped to a few seconds instead of the whole
|
||||
// process lifetime.
|
||||
func (ps *Plugins) Close() {
|
||||
if ps == nil {
|
||||
return
|
||||
}
|
||||
ps.save.Close()
|
||||
}
|
||||
|
||||
// LoadDir loads every .lua file in dir as a plugin. Files are loaded in
|
||||
@ -423,6 +598,12 @@ func (ps *Plugins) LoadSource(name, code string) error {
|
||||
p.UI = ui
|
||||
}
|
||||
|
||||
// Restore persisted state AFTER the plugin compiled and registered its
|
||||
// hooks, so the restore overwrites the compiled-in defaults instead of being
|
||||
// overwritten by them. Doing it earlier would mean a restart resets the
|
||||
// totals back to whatever the .lua source initialises them to.
|
||||
ps.restore(p)
|
||||
|
||||
ps.append(p)
|
||||
// Rebuild here rather than only in LoadDir: LoadSource is also the single-
|
||||
// plugin entry point (the WebUI upload path), and a caller that loads one
|
||||
@ -785,7 +966,21 @@ func (ps *Plugins) SetState(name string, state interface{}) error {
|
||||
// replacing it with an empty table would silently erase every
|
||||
// accumulated total, so the next request would start from zero
|
||||
// and the dashboard would show a sudden drop in spend.
|
||||
L.SetTop(0)
|
||||
//
|
||||
// The prices still have to be persisted here: returning before
|
||||
// the write below would leave the new price table in memory
|
||||
// only, and a restart would silently revert to the old prices
|
||||
// while the operator believed the change took effect.
|
||||
var curState, curPrices interface{}
|
||||
curState = snapshotField(L, "state")
|
||||
curPrices = snapshotField(L, "prices")
|
||||
if ps.dir != "" && !ps.stateDirDisabled && curState != nil {
|
||||
writeJSONAtomic(ps.stateFile(name), map[string]interface{}{
|
||||
"version": 1,
|
||||
"state": curState,
|
||||
"prices": curPrices,
|
||||
})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
body = rest
|
||||
@ -795,9 +990,218 @@ func (ps *Plugins) SetState(name string, state interface{}) error {
|
||||
pushGoValue(L, body)
|
||||
L.SetField(plug, "state")
|
||||
L.SetTop(0)
|
||||
// A state replacement is exactly the kind of thing an operator restarts the
|
||||
// gateway for, so it must survive the restart. The snapshot is taken from
|
||||
// the values just written — the caller already holds p.mu, so calling
|
||||
// persistNow here would deadlock on that same non-reentrant mutex.
|
||||
writeJSONAtomic(ps.stateFile(name), map[string]interface{}{
|
||||
"version": 1,
|
||||
"state": body,
|
||||
"prices": readPricesFromPayload(state),
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
// readPricesFromPayload recovers the prices table a caller passed to SetState,
|
||||
// which SetState stores on the plugin (not in state). Used only for the
|
||||
// persistence record, so that a restart restores configuration and history to
|
||||
// the two fields they belong in rather than collapsing them.
|
||||
func readPricesFromPayload(state interface{}) interface{} {
|
||||
if m, ok := state.(map[string]interface{}); ok {
|
||||
return m["prices"]
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// markDirty snapshots a plugin's state and hands it to the saver.
|
||||
//
|
||||
// It is called from the hook path, so it must not block on file I/O — that is
|
||||
// the saver's job. It DOES walk the Lua tables, because that has to happen
|
||||
// while the caller still holds p.mu and the VM is guaranteed alive; the flush
|
||||
// goroutine only ever sees the resulting Go values (see pendingState).
|
||||
//
|
||||
// Cost is one JSON conversion per request, which the billing plugin would pay
|
||||
// anyway inside its own hook.
|
||||
func (ps *Plugins) markDirtyLocked(p *Plugin) {
|
||||
if ps == nil || p == nil || ps.stateDirDisabled || p.state == nil || ps.dir == "" {
|
||||
return
|
||||
}
|
||||
// Caller already holds p.mu — taking it again would self-deadlock. That is
|
||||
// not hypothetical: the first version had markDirty lock p.mu and was called
|
||||
// from invoke(), which holds p.mu for the whole Lua call, so every hook call
|
||||
// deadlocked and the lua package's tests hung until the timeout.
|
||||
state, prices := readStateAndPrices(p.state.L)
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
ps.save.mark(pendingState{name: p.Info.Name, state: state, prices: prices})
|
||||
}
|
||||
|
||||
// ---------- state persistence ----------
|
||||
//
|
||||
// A plugin's state lives in the Lua VM, which dies with the process. For the
|
||||
// billing plugin that means a restart silently zeroes every total — measured on
|
||||
// this gateway: 218,241 requests and 26.2M prompt tokens gone after one
|
||||
// `systemctl restart`. For a REPORTING plugin whose whole purpose is the number
|
||||
// it accumulates, that is not a rounding error, it is the feature not working.
|
||||
//
|
||||
// Three rules shape this:
|
||||
//
|
||||
// 1. `prices` (configuration) and `state` (accumulated history) are persisted
|
||||
// to SEPARATE files. Restoring them together would let a price edit look
|
||||
// like a state reset, or a state restore resurrect stale prices — SetState
|
||||
// already keeps them apart in memory, and disk has to agree.
|
||||
// 2. Writes are atomic (temp file + rename). A crash mid-write must leave the
|
||||
// previous state readable, not a truncated JSON file that fails to parse on
|
||||
// the next boot and loses the total anyway.
|
||||
// 3. A corrupt or unreadable state file is a WARNING, never a startup error.
|
||||
// Forwarding must not depend on a plugin's bookkeeping surviving.
|
||||
|
||||
// stateFile is the per-plugin state path, kept beside the plugin source so an
|
||||
// operator can find (and delete) it next to the plugin it belongs to.
|
||||
func (ps *Plugins) stateFile(name string) string {
|
||||
return filepath.Join(ps.dir, "."+name+".state.json")
|
||||
}
|
||||
|
||||
// persistNow writes a plugin's state and prices to disk synchronously. Callers
|
||||
// on the hot path must use ps.markDirtyLocked instead; this is the flush worker and
|
||||
// the admin-state path, where durability matters more than latency.
|
||||
//
|
||||
// It is called after every
|
||||
// state mutation, so it must be cheap enough not to matter: the billing plugin
|
||||
// mutates on every request, and writing the whole state per request would put a
|
||||
// file write on the hot path.
|
||||
func (ps *Plugins) persistNow(p *Plugin) {
|
||||
if ps == nil || p == nil || ps.stateDirDisabled || p.state == nil || ps.dir == "" {
|
||||
return
|
||||
}
|
||||
p.mu.Lock()
|
||||
state, prices := readStateAndPrices(p.state.L)
|
||||
p.mu.Unlock()
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
writeJSONAtomic(ps.stateFile(p.Info.Name), map[string]interface{}{
|
||||
"version": 1,
|
||||
"state": state,
|
||||
"prices": prices,
|
||||
})
|
||||
}
|
||||
|
||||
// readStateAndPrices pulls both tables out of the Lua state. The caller holds
|
||||
// p.mu. Returns nils when the plugin table is missing, which happens for a
|
||||
// plugin that failed to compile.
|
||||
//
|
||||
// Each field is read by REBUILDING the stack from the global, because the
|
||||
// conversion helper (luaToJSON) ends with L.SetTop(0) on every path — it treats
|
||||
// the whole stack as its own. Reusing a saved index across two conversions
|
||||
// addresses a slot that no longer exists, and this binding aborts the process
|
||||
// (SIGABRT) instead of reporting a bad index. The first version did exactly
|
||||
// that and crashed inside the Lua C layer on every hook call.
|
||||
func readStateAndPrices(L *golua.State) (state, prices interface{}) {
|
||||
state = snapshotField(L, "state")
|
||||
prices = snapshotField(L, "prices")
|
||||
L.SetTop(0)
|
||||
return state, prices
|
||||
}
|
||||
|
||||
// snapshotField converts plugin.<field> into Go values, leaving the stack clean.
|
||||
// Caller holds p.mu; caller is responsible for the Lua state being alive.
|
||||
func snapshotField(L *golua.State, field string) interface{} {
|
||||
L.SetTop(0)
|
||||
defer L.SetTop(0)
|
||||
L.GetGlobal(pluginGlobal)
|
||||
if L.IsNil(-1) {
|
||||
return nil
|
||||
}
|
||||
L.GetField(L.GetTop(), field)
|
||||
if L.Type(-1) != golua.LUA_TTABLE {
|
||||
return nil
|
||||
}
|
||||
var out interface{}
|
||||
if err := luaToJSON(L, -1, &out); err != nil {
|
||||
return nil
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// writeJSONAtomic writes v to path via a temp file + rename, so a reader (or a
|
||||
// crash) never observes a half-written file.
|
||||
func writeJSONAtomic(path string, v interface{}) {
|
||||
b, err := json.MarshalIndent(v, "", " ")
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
dir := filepath.Dir(path)
|
||||
if err := os.MkdirAll(dir, 0755); err != nil {
|
||||
return
|
||||
}
|
||||
tmp, err := os.CreateTemp(dir, filepath.Base(path)+".tmp*")
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
if _, err := tmp.Write(b); err != nil {
|
||||
tmp.Close()
|
||||
os.Remove(tmpName)
|
||||
return
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
os.Remove(tmpName)
|
||||
return
|
||||
}
|
||||
if err := os.Rename(tmpName, path); err != nil {
|
||||
os.Remove(tmpName)
|
||||
}
|
||||
}
|
||||
|
||||
// restore re-applies a persisted state (and prices) onto a freshly compiled
|
||||
// plugin. Called once per plugin after it compiles and registers its hooks.
|
||||
//
|
||||
// Ordering matters: this runs AFTER compilation so the plugin's own defaults
|
||||
// exist, and it OVERWRITES them, so a restart continues from the saved totals
|
||||
// rather than from the values baked into the .lua source.
|
||||
func (ps *Plugins) restore(p *Plugin) {
|
||||
if ps == nil || p == nil || p.state == nil || ps.stateDirDisabled {
|
||||
return
|
||||
}
|
||||
b, err := os.ReadFile(ps.stateFile(p.Info.Name))
|
||||
if err != nil {
|
||||
return // no saved state yet: the compiled-in default stands
|
||||
}
|
||||
var saved struct {
|
||||
Version int `json:"version"`
|
||||
State interface{} `json:"state"`
|
||||
Prices interface{} `json:"prices"`
|
||||
}
|
||||
if err := json.Unmarshal(b, &saved); err != nil {
|
||||
// A corrupt file must not stop the gateway: the plugin keeps its
|
||||
// compiled-in defaults and the operator sees a warning in the log.
|
||||
ps.logf("plugin %s: ignoring unreadable state file %s: %v",
|
||||
p.Info.Name, ps.stateFile(p.Info.Name), err)
|
||||
return
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
L := p.state.L
|
||||
L.SetTop(0)
|
||||
defer L.SetTop(0)
|
||||
L.GetGlobal(pluginGlobal)
|
||||
if L.IsNil(-1) {
|
||||
return
|
||||
}
|
||||
plug := L.GetTop()
|
||||
if saved.Prices != nil {
|
||||
pushGoValue(L, saved.Prices)
|
||||
L.SetField(plug, "prices")
|
||||
}
|
||||
if saved.State != nil {
|
||||
pushGoValue(L, saved.State)
|
||||
L.SetField(plug, "state")
|
||||
}
|
||||
p.persisted = true
|
||||
}
|
||||
|
||||
// SetEnabled turns a plugin's dispatch on or off without touching its file.
|
||||
//
|
||||
// The state is on the Plugin record (not derived from disk) so a disable survives
|
||||
@ -960,12 +1364,26 @@ func (ps *Plugins) invoke(p *Plugin, fn string, payload map[string]interface{})
|
||||
if err := L.Call(1, 1); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Read the return value FIRST, then snapshot state for persistence.
|
||||
//
|
||||
// The order is load-bearing. markDirtyLocked walks the Lua tables and resets
|
||||
// the stack, so calling it before the return value is read destroyed the
|
||||
// hook's answer — every plugin that returned a table silently became a
|
||||
// plugin that "had no opinion". Reading first costs nothing and keeps the
|
||||
// documented merge contract intact.
|
||||
var out map[string]interface{}
|
||||
if L.GetTop() < 1 || L.IsNil(-1) {
|
||||
ps.markDirtyLocked(p)
|
||||
return nil, nil
|
||||
}
|
||||
var out map[string]interface{}
|
||||
if err := luaToJSON(L, -1, &out); err != nil {
|
||||
ps.markDirtyLocked(p)
|
||||
return nil, nil // not a table: treat as "no opinion"
|
||||
}
|
||||
// A hook that ran at all may have mutated plugin.state, whether or not it
|
||||
// returned anything. Marking here (not only on a returned table) is what
|
||||
// makes an accumulating plugin like billing durable: its totals change on
|
||||
// every call and it returns nil every time.
|
||||
ps.markDirtyLocked(p)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user