From b4006212e7c6c97e48bd65ac8925020982b5090a Mon Sep 17 00:00:00 2001 From: JianFeeeee Date: Sat, 15 Aug 2026 18:02:50 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20Lua=20=E6=8F=92=E4=BB=B6=20stage=20?= =?UTF-8?q?=E5=86=99=E5=9B=9E=E6=94=AF=E6=8C=81(=E4=B8=8E=20C=20ABI=20v2?= =?UTF-8?q?=20=E5=AF=B9=E9=BD=90)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit makeStageHandler 现在把 ctx 以 Lua table 引用传给 handler, handler 修改的字段(raw_message/llm_text/final_text/response/tool_calls/ tool_results/user_id/group_id/no_memory)同步写回内核 StageContext。 新增 TestLuaStageWriteback 验证 on_input/post_action 的写回。 --- internal/plugin/lua_plugin.go | 54 ++++++++++++++++++++++- internal/plugin/lua_plugin_test.go | 71 ++++++++++++++++++++++++++++++ 2 files changed, 124 insertions(+), 1 deletion(-) diff --git a/internal/plugin/lua_plugin.go b/internal/plugin/lua_plugin.go index 9718392..f724967 100644 --- a/internal/plugin/lua_plugin.go +++ b/internal/plugin/lua_plugin.go @@ -741,15 +741,67 @@ func makeStageHandler(plg *luaPlugin, stage sdk.Stage, fn *lua.LFunction) sdk.St if len(sc.ToolResults) > 0 { ctx["tool_results"] = jsonToIface(sc.ToolResults) } + ctxTbl := goValueToLua(L, ctx).(*lua.LTable) L.Push(fn) - L.Push(goValueToLua(L, ctx)) + L.Push(ctxTbl) if err := L.PCall(1, 0, nil); err != nil { return fmt.Errorf("lua stage %s: %w", stage, err) } + // 写回:Lua handler 对 ctx table 的字段修改同步回内核 StageContext + applyLuaStageResult(sc, luaValueToGo(ctxTbl)) return nil } } +// applyLuaStageResult 将 Lua stage handler 修改后的 ctx 字段写回内核 StageContext。 +// Lua 侧修改的字段以 Lua table(引用)形式读回,仅回写插件有权改写的键。 +func applyLuaStageResult(sc *sdk.StageContext, modified interface{}) { + m, ok := modified.(map[string]interface{}) + if !ok { + return + } + sc.Lock() + defer sc.Unlock() + if v, ok := m["raw_message"].(string); ok { + sc.RawMessage = v + } + if v, ok := m["llm_text"].(string); ok { + sc.LLMText = v + } + if v, ok := m["final_text"].(string); ok { + sc.FinalText = v + } + if v, ok := m["user_id"].(string); ok { + sc.UserID = v + } + if v, ok := m["group_id"].(string); ok { + sc.GroupID = v + } + if v, ok := m["no_memory"].(bool); ok { + sc.NoMemory = v + } + if v, ok := m["response"].(string); ok { + vv := v + sc.Response = &vv + } + if v, ok := m["tool_calls"].([]interface{}); ok && len(v) > 0 { + if b, err := json.Marshal(v); err == nil { + var tcs []sdk.ToolCall + if json.Unmarshal(b, &tcs) == nil { + sc.ToolCalls = tcs + } + } + } + if v, ok := m["tool_results"].([]interface{}); ok && len(v) > 0 { + if b, err := json.Marshal(v); err == nil { + var trs []sdk.ToolResult + if json.Unmarshal(b, &trs) == nil { + sc.ToolResults = trs + } + } + } +} + func parseStageScope(L *lua.LState) sdk.StageScope { if L.GetTop() >= 3 && L.ToString(3) == "own_tools" { return sdk.StageScopeOwnTools diff --git a/internal/plugin/lua_plugin_test.go b/internal/plugin/lua_plugin_test.go index 1827f52..25e4a42 100644 --- a/internal/plugin/lua_plugin_test.go +++ b/internal/plugin/lua_plugin_test.go @@ -526,3 +526,74 @@ return plugin t.Errorf("stage ctx raw_message = %v, want 'hi'", ctxTbl.RawGetString("raw_message")) } } + +func TestLuaStageWriteback(t *testing.T) { + dir := t.TempDir() + + os.WriteFile(filepath.Join(dir, "plugin.json"), []byte(`{"name":"wblua","entry":"main.lua"}`), 0644) + os.WriteFile(filepath.Join(dir, "main.lua"), []byte(` +local plugin = { name = "wblua" } + +function plugin.start(sdk) + sdk.register_stage("on_input", function(ctx) + ctx.raw_message = "[清洗]" .. ctx.raw_message + ctx.final_text = "改写后的最终文本" + end) + sdk.register_stage("post_action", function(ctx) + ctx.llm_text = ctx.llm_text .. "[尾部标记]" + end) +end + +function plugin.stop() end +return plugin +`), 0644) + + plg, err := tryLoadLua(dir, "wblua", nil) + if err != nil { + t.Fatalf("tryLoadLua failed: %v", err) + } + + var handlers = map[sdk.Stage]sdk.StageHandler{} + reg := internalConfig.NewConfigRegistry("") + sett := sdk.NewSettings("wblua", reg) + s := sdk.New("wblua", sdk.SDKConfig{ + Settings: sett, + RegStage: func(stage sdk.Stage, handler sdk.StageHandler) { + handlers[stage] = handler + }, + }) + + if err := plg.Start(s); err != nil { + t.Fatalf("Start failed: %v", err) + } + defer plg.Stop() + + // 触发 on_input stage,验证 Lua 修改写回内核 sc + onInput, ok := handlers[sdk.StageOnInput] + if !ok { + t.Fatal("on_input handler not registered") + } + sc := &sdk.StageContext{RawMessage: "原始消息"} + if err := onInput(sc); err != nil { + t.Fatalf("on_input: %v", err) + } + if sc.RawMessage != "[清洗]原始消息" { + t.Errorf("raw_message writeback: got %q, want %q", sc.RawMessage, "[清洗]原始消息") + } + if sc.FinalText != "改写后的最终文本" { + t.Errorf("final_text writeback: got %q", sc.FinalText) + } + + // 触发 post_action stage + post, ok := handlers[sdk.StagePostAction] + if !ok { + t.Fatal("post_action handler not registered") + } + sc2 := &sdk.StageContext{LLMText: "模型输出"} + if err := post(sc2); err != nil { + t.Fatalf("post_action: %v", err) + } + if sc2.LLMText != "模型输出[尾部标记]" { + t.Errorf("llm_text writeback: got %q, want %q", sc2.LLMText, "模型输出[尾部标记]") + } +}