mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-24 10:58:13 +00:00
feat: Lua 插件 stage 写回支持(与 C ABI v2 对齐)
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 的写回。
This commit is contained in:
@ -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
|
||||
|
||||
@ -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, "模型输出[尾部标记]")
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user