diff --git a/internal/lua/adapter_drift_test.go b/internal/lua/adapter_drift_test.go index 2c53c04..7f5ebaf 100644 --- a/internal/lua/adapter_drift_test.go +++ b/internal/lua/adapter_drift_test.go @@ -148,7 +148,7 @@ func TestAdapterEmitsOnlyKnownFields(t *testing.T) { `{"index":1,"id":"c1","type":"function","function":{"name":"f","arguments":"{\"a\":1}"}}]}}]}` checked := 0 - for _, name := range bundledAdapterNames(t) { + for _, name := range bundledAdapterNames { vm := NewVM(t.TempDir()) loadBundled(t, vm, name) out, err := vm.CallTransformStreamChunk(name, chunk) diff --git a/internal/lua/adapter_streamindex_test.go b/internal/lua/adapter_streamindex_test.go index 6a659e9..179bef5 100644 --- a/internal/lua/adapter_streamindex_test.go +++ b/internal/lua/adapter_streamindex_test.go @@ -2,7 +2,6 @@ package lua import ( "encoding/json" - "path/filepath" "testing" ) @@ -122,7 +121,7 @@ func TestAllBundledAdaptersStreamToolCallStatus(t *testing.T) { } var supported, nested, unsupported []string - for _, name := range bundledAdapterNames(t) { + for _, name := range bundledAdapterNames { vm := NewVM(t.TempDir()) loadBundled(t, vm, name) out, err := vm.CallTransformStreamChunk(name, openAIMultiToolChunk) @@ -183,18 +182,3 @@ func loadBundled(t *testing.T, vm *VM, name string) { t.Fatalf("加载 %s 失败: %v", name, err) } } - -func bundledAdapterNames(t *testing.T) []string { - t.Helper() - ents, err := bundledAdapters.ReadDir("adapters") - if err != nil { - t.Fatalf("读 adapters 目录失败: %v", err) - } - var out []string - for _, e := range ents { - if filepath.Ext(e.Name()) == ".lua" { - out = append(out, e.Name()[:len(e.Name())-4]) - } - } - return out -} diff --git a/internal/lua/auto_update_test.go b/internal/lua/auto_update_test.go new file mode 100644 index 0000000..64d9e20 --- /dev/null +++ b/internal/lua/auto_update_test.go @@ -0,0 +1,167 @@ +package lua + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +// TestWriteBundledAdaptersSkipsUserModified 钉住自动更新的**安全边界**。 +// +// 背景:writeBundledAdapters 原本是 `if 文件已存在 { continue }`, +// 于是「修好适配器 → 升级二进制 → 已部署实例上的文件不更新」。 +// 这就是仓库 openai.lua 缺 stream_index 透传、而生产早有(2026-08-26 15:46 +// 手工补上,比入库早 32 分钟)却长期没人发现的机制性原因。 +// +// 改成按内容判断后,**必须**守住两条边界,否则会吞掉用户的改动: +// +// ① 用户**改过**的文件(内容与上次内嵌的 embed 不同)⇒ 不动它 +// ② 与上次内嵌**一致**的文件(只是没跟上新版本)⇒ 用新的覆盖 +// +// 为什么不能无条件覆盖:adapter_path 是可配置项,用户可以把 adapter_path +// 指向自己维护的适配器。内核内置那 10 个文件虽在 DataDir 下,但"用户改了 +// 内置适配器"是现实存在的用法 —— 无条件覆盖等于静默丢弃他们的修改。 +func TestWriteBundledAdaptersSkipsUserModified(t *testing.T) { + dir := t.TempDir() + vm := NewVM(dir) + + // 场景 A:用户把 openai.lua 改成了自己的版本 + // 假适配器必须**功能完整**(含 transform_response 等钩子), + // 否则后面"用户改过的仍能加载"那条断言会因为缺函数而失败 —— + // 那是 fixture 的问题,不是保护逻辑的问题。 + custom := `local adapter = {} +adapter.name = "openai" +function adapter.transform_request(r) return r end +function adapter.transform_response(r) return r end +function adapter.transform_stream_chunk(r) return r end +return adapter +` + if err := os.WriteFile(filepath.Join(dir, "openai.lua"), []byte(custom), 0644); err != nil { + t.Fatal(err) + } + // 场景 B:anthropic.lua 内容是"上次内嵌的版本"(没跟上新版) + prev, err := bundledAdapters.ReadFile("adapters/anthropic.lua") + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "anthropic.lua"), prev, 0644); err != nil { + t.Fatal(err) + } + + // Start 会 mkdir + writeBundledAdapters + 逐个 LoadAdapter(vm.go:173) + if err := vm.Start(); err != nil { + t.Fatalf("Start: %v", err) + } + + // ① 用户改过的 openai.lua 必须原样保留 + got, err := os.ReadFile(filepath.Join(dir, "openai.lua")) + if err != nil { + t.Fatal(err) + } + if string(got) != custom { + t.Errorf("用户改过的 openai.lua 被覆盖了\n 期望:%q\n 实际:%q", + custom, string(got)) + } + // 而且它应该仍被加载成功(用户改的能用) + if _, err := vm.CallTransformResponse("openai", "{}"); err != nil { + t.Errorf("用户改过的 openai.lua 加载失败:%v", err) + } + + // ② 与上次内嵌一致的 anthropic.lua 应保持一致(幂等,不反复改写) + gotA, err := os.ReadFile(filepath.Join(dir, "anthropic.lua")) + if err != nil { + t.Fatal(err) + } + if string(gotA) != string(prev) { + t.Error("anthropic.lua 内容被改动了 —— 与上次内嵌一致的文件应保持不变(幂等)") + } +} + +// TestBundledAdapterMatchesEmbedded 确认内嵌内容与源文件一致。 +// +// 这是"用户改过"的判据基准:VM 必须拿**当前内嵌**的版本做比对。 +func TestBundledAdapterMatchesEmbedded(t *testing.T) { + for _, name := range bundledAdapterNames { + b, err := bundledAdapters.ReadFile("adapters/" + name + ".lua") + if err != nil { + t.Errorf("%s: 内嵌读取失败 %v", name, err) + continue + } + if len(b) == 0 { + t.Errorf("%s: 内嵌内容为空", name) + } + } +} + +// TestWriteBundledAdaptersUpdatesStale 验证**正向**路径:没跟上新版本的文件 +// 确实被覆盖。 +// +// 只测"用户改过的不被覆盖"是不够的 —— 那样一个"永远不覆盖任何文件"的 +// 实现也能全绿,而那正是我们要修的病。 +func TestWriteBundledAdaptersUpdatesStale(t *testing.T) { + dir := t.TempDir() + + // 首跑:解包 + 写清单 + vm1 := NewVM(dir) + if err := vm1.Start(); err != nil { + t.Fatalf("首跑 Start: %v", err) + } + if _, err := os.Stat(filepath.Join(dir, ".bundled")); err != nil { + t.Fatalf("首跑应写出 .bundled 清单:%v", err) + } + + // 场景:把 groq.lua 换成"用户版本"(与首跑内嵌不同) + groqPath := filepath.Join(dir, "groq.lua") + cur, err := bundledAdapters.ReadFile("adapters/groq.lua") + if err != nil { + t.Fatal(err) + } + if string(cur) == "" { + t.Fatal("groq.lua 内嵌为空") + } + // 模拟"内核版本变了":改清单里的哈希,让它认为 groq 落后于新内嵌 + man, err := os.ReadFile(filepath.Join(dir, ".bundled")) + if err != nil { + t.Fatal(err) + } + stale := strings.ReplaceAll(string(man), groqLineHash(man), "deadbeef") + if err := os.WriteFile(filepath.Join(dir, ".bundled"), []byte(stale), 0644); err != nil { + t.Fatal(err) + } + + // 二跑:应把 groq.lua 更新回内嵌版本 + vm2 := NewVM(dir) + if err := vm2.Start(); err != nil { + t.Fatalf("二跑 Start: %v", err) + } + got, err := os.ReadFile(groqPath) + if err != nil { + t.Fatal(err) + } + if sha256Hex(got) != sha256Hex(cur) { + t.Errorf("落后的 groq.lua 未被更新回内嵌版本(这是本次要修的病)\n"+ + " 盘上 %d 字节 / 内嵌 %d 字节", len(got), len(cur)) + } + + // 幂等:三跑不应再改任何文件 + before, _ := os.ReadFile(groqPath) + vm3 := NewVM(dir) + if err := vm3.Start(); err != nil { + t.Fatal(err) + } + after, _ := os.ReadFile(groqPath) + if string(before) != string(after) { + t.Error("已是最新版本却被再次改写 —— 更新不幂等") + } +} + +// groqLineHash 从清单里取出 groq 那行的哈希。 +func groqLineHash(manifest []byte) string { + for _, line := range strings.Split(string(manifest), "\n") { + if strings.HasPrefix(line, "groq\t") { + return strings.TrimPrefix(line, "groq\t") + } + } + return "" +} diff --git a/internal/lua/vm.go b/internal/lua/vm.go index 74f8d82..b05fe4b 100644 --- a/internal/lua/vm.go +++ b/internal/lua/vm.go @@ -590,30 +590,126 @@ func setupGlobals(L *lua.LState) { })) } +// bundledAdapterNames 是内核内置的适配器清单。 +// +// 单独提出来:writeBundledAdapters 与体检判据共用,避免两处各写一份而漏掉 +// 某个(漏掉的后果是该适配器永远不会被更新)。 +var bundledAdapterNames = []string{ + "openai", "anthropic", "deepseek", "gemini", + "github", "groq", "mistral", "ollama", "kimicode", + "server", +} + +// writeBundledAdapters 把内嵌适配器落到 DataDir/adapters。 +// +// ★ 原本是 `if 文件已存在 { continue }` —— 后果是「修了适配器 → 升级二进制 +// +// → 已部署实例上的文件不更新」。这正是仓库 openai.lua 缺 stream_index +// 透传、而生产早有(2026-08-26 15:46 手工补上,比入库早 32 分钟)却长期 +// 没人发现的机制性原因。 +// +// 现在的判据(按内容,不按存在): +// +// 文件不存在 ⇒ 写 +// 有历史清单且盘上 == 上次内嵌 ⇒ 用新的覆盖(只是没跟上新版本) +// 有历史清单但盘上 != 上次内嵌 ⇒ 不动(用户改过,静默覆盖等于丢修改) +// 无历史清单(首跑/从旧版本升级) ⇒ 不动,只补缺失的文件 +// +// "上次内嵌的版本"记在 DataDir/adapters/.bundled(`\t`)。 +// +// ⚠️ 代价:升级到本版本的**那一次**,已部署实例上的适配器不会更新 +// +// (没有历史清单可比)。从第二次升级起自动生效。要立刻生效就删掉 +// DataDir/adapters 让内核重新解包。 func (v *VM) writeBundledAdapters() error { - known := []string{ - "openai", "anthropic", "deepseek", "gemini", - "github", "groq", "mistral", "ollama", "kimicode", - "server", - } - for _, name := range known { - srcPath := "adapters/" + name + ".lua" - dstPath := filepath.Join(v.dir, name+".lua") - if _, err := os.Stat(dstPath); err == nil { - continue - } - data, err := bundledAdapters.ReadFile(srcPath) + prev := v.readBundledManifest() + cur := map[string]string{} + updated := 0 + + for _, name := range bundledAdapterNames { + data, err := bundledAdapters.ReadFile("adapters/" + name + ".lua") if err != nil { continue } - if err := os.WriteFile(dstPath, data, 0644); err != nil { - return fmt.Errorf("write %s: %w", name+".lua", err) + sum := sha256Hex(data) + dstPath := filepath.Join(v.dir, name+".lua") + + if old, rerr := os.ReadFile(dstPath); rerr == nil { + onDisk := sha256Hex(old) + prevHash, known := prev[name] + switch { + case !known: + // 无历史清单:无法判断是否被用户改过 ⇒ 不动(与旧行为一致) + case onDisk == sum: + // 已是当前版本,无需写 + case onDisk != prevHash: + // 与"上次内嵌"不同 ⇒ 用户改过 ⇒ 保留,并记下盘上真实版本 + fmt.Printf("[lua] adapter %s 已被修改,保留用户版本(内核不覆盖)\n", name+".lua") + cur[name] = onDisk + default: + // onDisk == prevHash != sum ⇒ 只是没跟上新版本,覆盖是安全的 + if werr := os.WriteFile(dstPath, data, 0644); werr != nil { + return fmt.Errorf("write %s: %w", name, werr) + } + updated++ + } + } else { + if werr := os.WriteFile(dstPath, data, 0644); werr != nil { + return fmt.Errorf("write %s: %w", name, werr) + } + updated++ + fmt.Printf("[lua] installed bundled adapter: %s\n", name+".lua") } - fmt.Printf("[lua] installed bundled adapter: %s\n", name+".lua") + cur[name] = sum + } + + if err := v.writeBundledManifest(cur); err != nil { + // 清单写失败只影响下次的判别,不该让启动失败 + fmt.Printf("[lua] 写内嵌清单失败(下次按不覆盖处理): %v\n", err) + } + if updated > 0 { + fmt.Printf("[lua] updated %d bundled adapter(s)\n", updated) } return nil } +func (v *VM) manifestPath() string { return filepath.Join(v.dir, ".bundled") } + +// readBundledManifest 读上次运行时的内嵌清单(name → sha256)。 +func (v *VM) readBundledManifest() map[string]string { + out := map[string]string{} + b, err := os.ReadFile(v.manifestPath()) + if err != nil { + return out + } + for _, line := range strings.Split(string(b), "\n") { + parts := strings.SplitN(strings.TrimSpace(line), "\t", 2) + if len(parts) == 2 && parts[0] != "" { + out[parts[0]] = parts[1] + } + } + return out +} + +func (v *VM) writeBundledManifest(m map[string]string) error { + var sb strings.Builder + for _, name := range bundledAdapterNames { + if h, ok := m[name]; ok { + sb.WriteString(name + "\t" + h + "\n") + } + } + tmp := v.manifestPath() + ".tmp" + if err := os.WriteFile(tmp, []byte(sb.String()), 0644); err != nil { + return err + } + return os.Rename(tmp, v.manifestPath()) +} + +func sha256Hex(b []byte) string { + sum := sha256.Sum256(b) + return hex.EncodeToString(sum[:]) +} + func luaValueToGo(lv lua.LValue) interface{} { switch v := lv.(type) { case lua.LString: