mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-03 15:53:56 +00:00
feat(memory): 新增可声明的召回轴 RecallPolicy(与 prune 正交)
问题:召回(把 L2/L3 相关记忆注入本轮)此前不可声明、也不受任何 SDK 字段 控制——它只在任务开始时对 f.Input 无条件跑一次。于是 qq_get_message 取回 真实正文后只触发 Prune(裁剪),从不触发召回;而中断通知的 meta 文本反而 会去召回(词不对题,命中一堆泛实体)。 改动: - SDK 新增 RecallPolicy(none|auto) 轴,落在 InjectOptions / ChannelDef / ToolDef 三个声明面,与 ContextPolicy 正交(裁剪 vs 召回)。默认值与 prune 刻意相反:输入/注入默认 auto(保持既有「每条输入都召回」), 工具默认 none(工具输出多为噪声,按需声明)。 - 内核:recallDeclared 按 注入点 > 通道 > 默认auto 解析;输入侧用它决定 是否注入记忆索引;工具侧 ContextPolicy/RecallPolicy 共用同一份清洗后 query,一次相关性过程分别 prune / recall;召回以 system 消息挂到消息 末尾(同任务内替换而非累加)。 - 管线:proc RPC(inject/register + 校验)、lua 键、io payload 全量透传。 - QQ 插件:中断与 qq 通道声明 RecallPolicy=none(meta 不是内容); qq_get_message 声明 RecallPolicy=auto(取回正文后据正文召回)。 测试:新增 recallpolicy_test.go(core)与 proc 校验用例; go build ./... 通过,go test ./internal/... 全通过,SDK 模块与 qq 插件测试通过。
This commit is contained in:
@ -405,6 +405,29 @@ func (a *Agent) pruneDeclared(evt *agentIO.InputEvent) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// recallDeclared 判定这次输入是否要触发记忆召回(注入)。
|
||||
//
|
||||
// 与 pruneDeclared **正交**:prune 管“踢出去”(归档低相关 L0 事件),
|
||||
// recall 管“取进来”(把 L2/L3 相关记忆注入本轮)。
|
||||
//
|
||||
// 默认值与 prune 刻意相反:召回是只读增量、日常对话本就需要,所以**默认 auto**;
|
||||
// 只有显式声明 recall_policy=none(如中断通知的 meta 文本)才关闭。
|
||||
// 优先级同 prune:注入点(payload)> 通道(ChannelDef)> 默认 auto。
|
||||
func (a *Agent) recallDeclared(evt *agentIO.InputEvent) bool {
|
||||
if evt == nil {
|
||||
return true
|
||||
}
|
||||
if p, ok := evt.Payload["recall_policy"].(string); ok && p != "" {
|
||||
return p != pubsdk.RecallPolicyNone
|
||||
}
|
||||
if a.io != nil {
|
||||
if chDef, ok := a.io.GetInputChannelDef(evt.Source); ok && chDef.RecallPolicy != "" {
|
||||
return chDef.RecallPolicy != pubsdk.RecallPolicyNone
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// cleanInputFor 解析这条输入在计算层应当使用的清洗文本。
|
||||
//
|
||||
// 优先级:注入点声明的 cleaner(payload.cleaner_name,引用某个已注册的通道
|
||||
|
||||
@ -139,6 +139,9 @@ func policySuffix(ch agentIO.InputChannel) string {
|
||||
if ch.Def.ContextPolicy != "" && ch.Def.ContextPolicy != "none" {
|
||||
m = append(m, "裁剪:"+ch.Def.ContextPolicy)
|
||||
}
|
||||
if ch.Def.RecallPolicy != "" && ch.Def.RecallPolicy != "auto" {
|
||||
m = append(m, "召回:"+ch.Def.RecallPolicy)
|
||||
}
|
||||
if len(m) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
@ -81,6 +81,29 @@ func (a *Agent) toolOutputForQuery(toolName, raw string) string {
|
||||
return raw
|
||||
}
|
||||
|
||||
// recallMsgMarker 是工具触发召回时注入的 system 消息前缀。
|
||||
// 用它做去重与替换的识别标(与用户/中断的 system 消息区分开)。
|
||||
const recallMsgMarker = "【记忆召回】"
|
||||
|
||||
// appendOrReplaceRecall 把一段召回文本作为 system 消息挂到消息末尾。
|
||||
//
|
||||
// 同一任务内多次触发(如模型多次调用 qq_get_message)时**替换**上一条召回,
|
||||
// 而不是累加:否则召回会线性叠进 prompt,把上下文与 token 预算越挤越紧。
|
||||
// 替换位置固定在末尾,不影响 tool/assistant 消息的配对。
|
||||
func appendOrReplaceRecall(msgs []agentAPI.Message, recallText string) []agentAPI.Message {
|
||||
if recallText == "" {
|
||||
return msgs
|
||||
}
|
||||
full := recallMsgMarker + "\n" + recallText
|
||||
for i := len(msgs) - 1; i >= 0; i-- {
|
||||
if msgs[i].Role == "system" && strings.HasPrefix(msgs[i].Content, recallMsgMarker) {
|
||||
msgs[i].Content = full
|
||||
return msgs
|
||||
}
|
||||
}
|
||||
return append(msgs, agentAPI.Message{Role: "system", Content: full})
|
||||
}
|
||||
|
||||
// dropContinuationPlaceholders 移除此前由本机制插入的 user 占位。
|
||||
//
|
||||
// 为什么必须移除而不仅仅是“不再追加”:`msgs` 在循环外创建、循环内只增不减,
|
||||
|
||||
141
internal/agent/core/recallpolicy_test.go
Normal file
141
internal/agent/core/recallpolicy_test.go
Normal file
@ -0,0 +1,141 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
agentAPI "gitcode.com/JianFeeeee/HomeAgent/internal/agent/api"
|
||||
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||
pubsdk "gitcode.com/JianFeeeee/homeagent-sdk/sdk"
|
||||
)
|
||||
|
||||
// 这一组测试锁死「默认召回、可显式关闭」这条语义。
|
||||
//
|
||||
// 与 prune 刻意相反:裁剪是破坏性的、默认关;召回是只读增量、默认开。
|
||||
// 两者正交,一根 ContextPolicy 表达不了 2×2 的组合(只召回不裁剪 / 只裁不召回)。
|
||||
func TestRecallDeclared_DefaultsToRecall(t *testing.T) {
|
||||
m := agentIO.NewIOManager()
|
||||
a := &Agent{io: m}
|
||||
|
||||
// 没有任何声明 → 默认召回(保持既有"每条输入都召回"的行为)。
|
||||
if !a.recallDeclared(&agentIO.InputEvent{Source: "unknown", Payload: map[string]interface{}{}}) {
|
||||
t.Fatal("未声明的输入默认必须召回")
|
||||
}
|
||||
// 通道注册了但没设 RecallPolicy → 仍默认召回。
|
||||
m.RegisterInputChannel("plain", pubsdk.ChannelDef{})
|
||||
if !a.recallDeclared(&agentIO.InputEvent{Source: "plain", Payload: map[string]interface{}{}}) {
|
||||
t.Fatal("ChannelDef 未设 RecallPolicy 应默认召回")
|
||||
}
|
||||
// nil 事件不能 panic,且按默认召回。
|
||||
if !a.recallDeclared(nil) {
|
||||
t.Fatal("nil 事件应默认召回")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecallDeclared_ChannelOptOut(t *testing.T) {
|
||||
m := agentIO.NewIOManager()
|
||||
m.RegisterInputChannel("meta", pubsdk.ChannelDef{RecallPolicy: pubsdk.RecallPolicyNone})
|
||||
m.RegisterInputChannel("talk", pubsdk.ChannelDef{RecallPolicy: pubsdk.RecallPolicyAuto})
|
||||
a := &Agent{io: m}
|
||||
|
||||
if a.recallDeclared(&agentIO.InputEvent{Source: "meta", Payload: map[string]interface{}{}}) {
|
||||
t.Fatal("通道声明 none 不应召回")
|
||||
}
|
||||
if !a.recallDeclared(&agentIO.InputEvent{Source: "talk", Payload: map[string]interface{}{}}) {
|
||||
t.Fatal("通道声明 auto 应召回")
|
||||
}
|
||||
}
|
||||
|
||||
// 注入点声明优先于通道定义:同一通道下的不同注入可以有不同意图。
|
||||
func TestRecallDeclared_InjectionOverridesChannel(t *testing.T) {
|
||||
m := agentIO.NewIOManager()
|
||||
a := &Agent{io: m}
|
||||
m.RegisterInputChannel("qq", pubsdk.ChannelDef{RecallPolicy: pubsdk.RecallPolicyNone})
|
||||
|
||||
evt := &agentIO.InputEvent{Source: "qq", Payload: map[string]interface{}{
|
||||
"recall_policy": pubsdk.RecallPolicyAuto,
|
||||
}}
|
||||
if !a.recallDeclared(evt) {
|
||||
t.Fatal("注入点声明 auto 应覆盖通道的 none")
|
||||
}
|
||||
|
||||
m.RegisterInputChannel("plain", pubsdk.ChannelDef{RecallPolicy: pubsdk.RecallPolicyAuto})
|
||||
evt = &agentIO.InputEvent{Source: "plain", Payload: map[string]interface{}{
|
||||
"recall_policy": pubsdk.RecallPolicyNone,
|
||||
}}
|
||||
if a.recallDeclared(evt) {
|
||||
t.Fatal("注入点声明 none 应覆盖通道的 auto")
|
||||
}
|
||||
}
|
||||
|
||||
// buildTaskMemoryContext 在声明 none 时必须返回空串(不注入记忆索引)。
|
||||
func TestBuildTaskMemoryContext_RespectsPolicy(t *testing.T) {
|
||||
m := agentIO.NewIOManager()
|
||||
m.RegisterInputChannel("meta", pubsdk.ChannelDef{RecallPolicy: pubsdk.RecallPolicyNone})
|
||||
a := &Agent{io: m, indexer: newTestIndexer(t, "咖啡", "张三")}
|
||||
|
||||
f := &TaskFrame{Evt: &agentIO.InputEvent{Source: "meta", Payload: map[string]interface{}{}}}
|
||||
if got := a.buildTaskMemoryContext(f, "咖啡", 0); got != "" {
|
||||
t.Fatalf("声明 none 时不应注入记忆,实际 %q", got)
|
||||
}
|
||||
|
||||
f2 := &TaskFrame{Evt: &agentIO.InputEvent{Source: "plain", Payload: map[string]interface{}{}}}
|
||||
if got := a.buildTaskMemoryContext(f2, "咖啡", 0); !strings.Contains(got, "【记忆索引】") {
|
||||
t.Fatalf("默认应注入记忆索引,实际 %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// 工具触发的召回:以(清洗后的)工具输出为 query,产出可注入的记忆文本。
|
||||
func TestRecallTextFor_UsesQuery(t *testing.T) {
|
||||
a := &Agent{indexer: newTestIndexer(t, "咖啡", "张三")}
|
||||
got := a.recallTextFor("咖啡", "tool:test")
|
||||
if !strings.Contains(got, "【记忆索引】") {
|
||||
t.Fatalf("应产出记忆索引文本,实际 %q", got)
|
||||
}
|
||||
// 空 query 或无 indexer 时不产出、不 panic。
|
||||
if got := a.recallTextFor("", "tool:test"); got != "" {
|
||||
t.Fatalf("空 query 应返回空串,实际 %q", got)
|
||||
}
|
||||
if got := (&Agent{}).recallTextFor("咖啡", "tool:test"); got != "" {
|
||||
t.Fatalf("无 indexer 应返回空串,实际 %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// 召回文本以 system 消息挂在末尾;同一任务内多次触发是**替换**而非累加。
|
||||
func TestAppendOrReplaceRecall(t *testing.T) {
|
||||
msgs := []agentAPI.Message{{Role: "user", Content: "hi"}}
|
||||
msgs = appendOrReplaceRecall(msgs, "第一段")
|
||||
if len(msgs) != 2 || msgs[1].Role != "system" || !strings.Contains(msgs[1].Content, "第一段") {
|
||||
t.Fatalf("首次应追加一条 system 召回消息,实际 %+v", msgs)
|
||||
}
|
||||
msgs = appendOrReplaceRecall(msgs, "第二段")
|
||||
if len(msgs) != 2 {
|
||||
t.Fatalf("再次触发应替换而非累加,实际 %d 条", len(msgs))
|
||||
}
|
||||
if !strings.Contains(msgs[1].Content, "第二段") || strings.Contains(msgs[1].Content, "第一段") {
|
||||
t.Fatalf("替换后应只含最新召回,实际 %q", msgs[1].Content)
|
||||
}
|
||||
if msgs = appendOrReplaceRecall(msgs, ""); len(msgs) != 2 {
|
||||
t.Fatalf("空召回不应改变消息,实际 %d 条", len(msgs))
|
||||
}
|
||||
}
|
||||
|
||||
// newTestIndexer 造一个只含给定实体的图记忆 + 已同步的索引器。
|
||||
func newTestIndexer(t *testing.T, subject, object string) *memory.Indexer {
|
||||
t.Helper()
|
||||
db, err := memory.NewGraphDB(filepath.Join(t.TempDir(), "graph.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("NewGraphDB: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
if _, _, err := db.Commit([]memory.Triple{{Subject: subject, Relation: "喜欢", Object: object}}, "s", 0); err != nil {
|
||||
t.Fatalf("Commit: %v", err)
|
||||
}
|
||||
idx := memory.NewIndexer(db)
|
||||
if err := idx.Sync(); err != nil {
|
||||
t.Fatalf("Sync: %v", err)
|
||||
}
|
||||
return idx
|
||||
}
|
||||
@ -275,7 +275,7 @@ func (a *Agent) rebaseFramePrefix(f *TaskFrame) {
|
||||
tail := append([]agentAPI.Message(nil), f.Msgs[f.PrefixLen:]...)
|
||||
|
||||
budget := ComputeTokenBudget(a.provider, a.systemPrompt)
|
||||
memContext := a.buildMemoryContext(f.Input, budget.MemoryTokens)
|
||||
memContext := a.buildTaskMemoryContext(f, f.Input, budget.MemoryTokens)
|
||||
sysPrompt := a.buildSystemPrompt(memContext, f.Input)
|
||||
prefix := a.buildMessages(sysPrompt, f.Input, a.contextTokenBudget(budget))
|
||||
|
||||
@ -497,7 +497,7 @@ func (a *Agent) step(f *TaskFrame) stepOutcome {
|
||||
func (a *Agent) stepPrepare(f *TaskFrame) stepOutcome {
|
||||
budget := ComputeTokenBudget(a.provider, a.systemPrompt)
|
||||
|
||||
memContext := a.buildMemoryContext(f.Input, budget.MemoryTokens)
|
||||
memContext := a.buildTaskMemoryContext(f, f.Input, budget.MemoryTokens)
|
||||
sysPrompt := a.buildSystemPrompt(memContext, f.Input)
|
||||
f.Tools = a.buildToolDefs()
|
||||
|
||||
@ -743,16 +743,27 @@ func (a *Agent) stepToolAfter(f *TaskFrame) stepOutcome {
|
||||
result = r
|
||||
}
|
||||
}
|
||||
// ContextPolicy: prune 工具调用后执行上下文裁剪(§13.8)
|
||||
if def := a.stageHost.ToolDef(tc.Name); def != nil && def.ContextPolicy == "prune" {
|
||||
if a.context != nil {
|
||||
topK := a.maxContextSize - 1
|
||||
if topK < 1 {
|
||||
topK = 1
|
||||
// 工具后处理:一次相关性过程,两个**正交**声明——
|
||||
// ContextPolicy=prune → 裁剪(踢出去,归档低相关 L0 事件)
|
||||
// RecallPolicy=auto → 召回(取进来,注入 L2/L3 相关记忆)
|
||||
// 两者共用同一份**清洗后**的 query:查询向量取清洗后的有效内容,否则噪声
|
||||
// (ANSI/base64/JSON 包装)会把相关性打分带偏,裁错事件、召回错记忆。
|
||||
var recallText string
|
||||
if def := a.stageHost.ToolDef(tc.Name); def != nil {
|
||||
needPrune := def.ContextPolicy == sdk.ContextPolicyPrune
|
||||
needRecall := def.RecallPolicy == sdk.RecallPolicyAuto
|
||||
if needPrune || needRecall {
|
||||
query := a.toolOutputForQuery(tc.Name, result)
|
||||
if needPrune && a.context != nil {
|
||||
topK := a.maxContextSize - 1
|
||||
if topK < 1 {
|
||||
topK = 1
|
||||
}
|
||||
a.context.Prune(query, topK, a.docStore)
|
||||
}
|
||||
if needRecall {
|
||||
recallText = a.recallTextFor(query, "tool:"+tc.Name)
|
||||
}
|
||||
// 查询向量取**清洗后**的有效内容,否则噪声(ANSI/base64/JSON 包装)
|
||||
// 会把相关性打分带偏,裁掉本该保留的事件。
|
||||
a.context.Prune(a.toolOutputForQuery(tc.Name, result), topK, a.docStore)
|
||||
}
|
||||
}
|
||||
|
||||
@ -822,6 +833,10 @@ func (a *Agent) stepToolAfter(f *TaskFrame) stepOutcome {
|
||||
// 必须紧跟在 toolMsg 之后:中间插入其他消息会让 tool_call_id 配对断开。
|
||||
f.Msgs = append(f.Msgs, *mediaMsg)
|
||||
}
|
||||
// 召回作为 system 消息挂在末尾(tool/assistant 配对已完成,插入此处不断链)。
|
||||
if recallText != "" {
|
||||
f.Msgs = appendOrReplaceRecall(f.Msgs, recallText)
|
||||
}
|
||||
|
||||
a.publishEvent(events.EventToolCall, map[string]interface{}{
|
||||
"tool": tc.Name,
|
||||
|
||||
@ -2,6 +2,7 @@ package core
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
|
||||
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
|
||||
@ -39,6 +40,38 @@ func (a *Agent) buildMemoryContext(input string, maxTokens int) string {
|
||||
return s
|
||||
}
|
||||
|
||||
// buildTaskMemoryContext 按本任务声明的召回策略决定是否注入记忆索引。
|
||||
//
|
||||
// 默认 auto(保持“每条输入都召回”的既有行为);输入/注入声明
|
||||
// recall_policy=none 时返回空串,从而不注入记忆。策略与裁剪(ContextPolicy)正交。
|
||||
func (a *Agent) buildTaskMemoryContext(f *TaskFrame, input string, maxTokens int) string {
|
||||
if f != nil && !a.recallDeclared(f.Evt) {
|
||||
return ""
|
||||
}
|
||||
return a.buildMemoryContext(input, maxTokens)
|
||||
}
|
||||
|
||||
// recallTextFor 以 query 触发一次记忆召回,返回可注入的文本(空串表示无)。
|
||||
//
|
||||
// 这是“召回”侧的单一入口:与 Prune 共用同一份**清洗后**的 query,
|
||||
// 使“取进来”(召回)与“踢出去”(裁剪)落在同一个相关性过程上。
|
||||
// trigger 仅用于日志溯源(如 "tool:qq_get_message")。
|
||||
func (a *Agent) recallTextFor(query, trigger string) string {
|
||||
if query == "" || a.indexer == nil {
|
||||
return ""
|
||||
}
|
||||
memTokens := 0 // 0 = 不截断
|
||||
if a.provider != nil {
|
||||
memTokens = ComputeTokenBudget(a.provider, a.systemPrompt).MemoryTokens
|
||||
}
|
||||
text := a.buildMemoryContext(query, memTokens)
|
||||
if text == "" {
|
||||
return ""
|
||||
}
|
||||
log.Printf("[agent] memory recall (%s): injected %d chars", trigger, len(text))
|
||||
return text
|
||||
}
|
||||
|
||||
// expandPromptVars 展开自定义提示词(人格卡)里的版本占位符。
|
||||
//
|
||||
// 为什么需要:人格卡是**配置项**,一旦写死版本号就会随内核发版而说谎 ——
|
||||
|
||||
@ -352,6 +352,9 @@ func applyInjectOpts(payload map[string]interface{}, opts InjectOptions) {
|
||||
if opts.ContextPolicy != "" {
|
||||
payload["context_policy"] = opts.ContextPolicy
|
||||
}
|
||||
if opts.RecallPolicy != "" {
|
||||
payload["recall_policy"] = opts.RecallPolicy
|
||||
}
|
||||
if opts.CleanerName != "" {
|
||||
payload["cleaner_name"] = opts.CleanerName
|
||||
}
|
||||
|
||||
@ -47,6 +47,7 @@ func TestInjectTextOpts_CarriesFlags(t *testing.T) {
|
||||
m.InjectTextOpts("src", "chan", "hello", InjectOptions{
|
||||
NoMemory: true,
|
||||
ContextPolicy: "prune",
|
||||
RecallPolicy: "none",
|
||||
CleanerName: "clean_me",
|
||||
})
|
||||
evt := drainOne(t, m.InputChan())
|
||||
@ -57,6 +58,9 @@ func TestInjectTextOpts_CarriesFlags(t *testing.T) {
|
||||
if evt.Payload["context_policy"] != "prune" {
|
||||
t.Errorf("context_policy 未传递: %v", evt.Payload["context_policy"])
|
||||
}
|
||||
if evt.Payload["recall_policy"] != "none" {
|
||||
t.Errorf("recall_policy 未传递: %v", evt.Payload["recall_policy"])
|
||||
}
|
||||
if evt.Payload["cleaner_name"] != "clean_me" {
|
||||
t.Errorf("cleaner_name 未传递: %v", evt.Payload["cleaner_name"])
|
||||
}
|
||||
@ -71,12 +75,15 @@ func TestInjectTextOpts_CarriesFlags(t *testing.T) {
|
||||
// 中断注入走另一条队列,标志位同样要带上(用户已确认中断允许声明 prune)。
|
||||
func TestInjectInterruptTextOpts_CarriesFlags(t *testing.T) {
|
||||
m := NewIOManager()
|
||||
m.InjectInterruptTextOpts("src", "chan", "alert", InjectOptions{ContextPolicy: "prune"})
|
||||
m.InjectInterruptTextOpts("src", "chan", "alert", InjectOptions{ContextPolicy: "prune", RecallPolicy: "none"})
|
||||
evt := drainOne(t, m.InputInterruptChan())
|
||||
|
||||
if evt.Payload["context_policy"] != "prune" {
|
||||
t.Errorf("中断注入的 context_policy 未传递: %v", evt.Payload)
|
||||
}
|
||||
if evt.Payload["recall_policy"] != "none" {
|
||||
t.Errorf("中断注入的 recall_policy 未传递: %v", evt.Payload)
|
||||
}
|
||||
if evt.Payload["type"] != "text" || evt.Payload["content"] != "alert" {
|
||||
t.Errorf("中断注入的基本字段不对: %v", evt.Payload)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user