mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-03 15:53:56 +00:00
feat: NoMemory/Cleaner memory system + doc update
- _sdk_local/ removed (moved to standalone sdk repo) - internal/agent/core: NoMemory/Cleaner data-flow breakpoints - internal/memory: clean_text, document store refactor - internal/plugin/registry.go: plugin API alignment - docs: PLUGIN_DEV.md, ARCHITECTURE.md NoMemory/Cleaner docs - plan.md, review.md: status update
This commit is contained in:
@ -16,13 +16,19 @@ import (
|
||||
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
|
||||
)
|
||||
|
||||
type ToolResultItem struct {
|
||||
Name string `json:"name"`
|
||||
Output string `json:"output"`
|
||||
}
|
||||
|
||||
type ContextEvent struct {
|
||||
Timestamp time.Time `json:"timestamp"`
|
||||
Source string `json:"source"`
|
||||
Input string `json:"input"`
|
||||
Response string `json:"response,omitempty"`
|
||||
ToolsUsed []string `json:"tools_used,omitempty"`
|
||||
Vector vector.Vector `json:"-"`
|
||||
Timestamp time.Time `json:"timestamp"`
|
||||
Source string `json:"source"`
|
||||
Input string `json:"input"`
|
||||
Response string `json:"response,omitempty"`
|
||||
ToolsUsed []string `json:"tools_used,omitempty"`
|
||||
ToolResults []ToolResultItem `json:"tool_results,omitempty"`
|
||||
Vector vector.Vector `json:"-"`
|
||||
}
|
||||
|
||||
const contextFlushInterval = 5 * time.Second
|
||||
@ -64,25 +70,49 @@ func (c *RelevanceContext) load() {
|
||||
return
|
||||
}
|
||||
for _, evt := range events {
|
||||
evt.Input = memory.CleanText(evt.Input)
|
||||
evt.Vector = c.computeVector(evt)
|
||||
}
|
||||
c.events = events
|
||||
}
|
||||
|
||||
func textForVector(evt *ContextEvent) string {
|
||||
func textForVector(evt *ContextEvent, toolDefLookup func(name string) *sdk.ToolDef) string {
|
||||
var text string
|
||||
switch {
|
||||
case evt.Source == "agent" && evt.Response != "":
|
||||
return memory.CleanText(evt.Response)
|
||||
text = evt.Response
|
||||
case evt.Source == "cold_storage":
|
||||
return memory.CleanText(evt.Input + " " + evt.Response)
|
||||
text = evt.Input + " " + evt.Response
|
||||
default:
|
||||
return memory.CleanText(evt.Input)
|
||||
text = evt.Input
|
||||
}
|
||||
|
||||
// 计算层:附加工具输出,NoMemory 跳过,其余经 Cleaner 过滤
|
||||
if toolDefLookup != nil {
|
||||
noMemory := make(map[string]bool)
|
||||
for _, tr := range evt.ToolResults {
|
||||
def := toolDefLookup(tr.Name)
|
||||
if def != nil && def.NoMemory {
|
||||
noMemory[tr.Name] = true
|
||||
}
|
||||
}
|
||||
for _, tr := range evt.ToolResults {
|
||||
if noMemory[tr.Name] {
|
||||
continue
|
||||
}
|
||||
cleaned := tr.Output
|
||||
def := toolDefLookup(tr.Name)
|
||||
if def != nil && def.Cleaner != nil {
|
||||
cleaned = def.Cleaner(cleaned)
|
||||
}
|
||||
text += " " + cleaned
|
||||
}
|
||||
}
|
||||
|
||||
return memory.CleanText(text)
|
||||
}
|
||||
|
||||
func (c *RelevanceContext) computeVector(evt *ContextEvent) vector.Vector {
|
||||
return c.embedder.Vectorize(textForVector(evt))
|
||||
return c.embedder.Vectorize(textForVector(evt, c.toolDefLookup))
|
||||
}
|
||||
|
||||
func (c *RelevanceContext) Save() error {
|
||||
@ -103,7 +133,6 @@ func (c *RelevanceContext) Append(evt ContextEvent) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
evt.Input = memory.CleanText(evt.Input)
|
||||
evt.Vector = c.computeVector(&evt)
|
||||
c.events = append(c.events, &evt)
|
||||
|
||||
@ -199,25 +228,19 @@ func (c *RelevanceContext) Prune(currentInput string, topK int, docStore *docume
|
||||
|
||||
archived := 0
|
||||
if docStore != nil && len(archive) > 0 {
|
||||
var filtered []scored
|
||||
for _, s := range archive {
|
||||
if hasNoMemoryTool(s.event.ToolsUsed, c.toolDefLookup) {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, s)
|
||||
}
|
||||
entries := make([]document.ContextEntry, len(filtered))
|
||||
for i, s := range filtered {
|
||||
entries := make([]document.ContextEntry, len(archive))
|
||||
for i, s := range archive {
|
||||
entries[i] = document.ContextEntry{
|
||||
Timestamp: s.event.Timestamp,
|
||||
Source: s.event.Source,
|
||||
Content: s.event.Input,
|
||||
Response: s.event.Response,
|
||||
Timestamp: s.event.Timestamp,
|
||||
Source: s.event.Source,
|
||||
Content: s.event.Input,
|
||||
Response: s.event.Response,
|
||||
ToolResults: convertToolResults(s.event.ToolResults),
|
||||
}
|
||||
}
|
||||
doc, err := docStore.ContextToDoc("context_archived", entries, c.embedder)
|
||||
if err == nil && doc != nil {
|
||||
archived = len(filtered)
|
||||
archived = len(entries)
|
||||
}
|
||||
}
|
||||
|
||||
@ -266,14 +289,15 @@ func (c *RelevanceContext) Len() int {
|
||||
return len(c.events)
|
||||
}
|
||||
|
||||
func hasNoMemoryTool(toolsUsed []string, lookup func(string) *sdk.ToolDef) bool {
|
||||
if lookup == nil {
|
||||
return false
|
||||
func convertToolResults(items []ToolResultItem) []document.ToolResultItem {
|
||||
if items == nil {
|
||||
return nil
|
||||
}
|
||||
for _, name := range toolsUsed {
|
||||
if def := lookup(name); def != nil && def.NoMemory {
|
||||
return true
|
||||
}
|
||||
result := make([]document.ToolResultItem, len(items))
|
||||
for i, item := range items {
|
||||
result[i] = document.ToolResultItem{Name: item.Name, Output: item.Output}
|
||||
}
|
||||
return false
|
||||
return result
|
||||
}
|
||||
|
||||
|
||||
|
||||
@ -6,7 +6,6 @@ import (
|
||||
"time"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
|
||||
)
|
||||
|
||||
func newTestCtx() *RelevanceContext {
|
||||
@ -238,40 +237,6 @@ func containsStr(s, substr string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func TestHasNoMemoryTool(t *testing.T) {
|
||||
host := NewStageHost()
|
||||
host.RegisterTool("no_mem_tool", sdk.ToolDef{Name: "no_mem_tool", NoMemory: true}, nil)
|
||||
host.RegisterTool("mem_tool", sdk.ToolDef{Name: "mem_tool"}, nil)
|
||||
lookup := host.ToolDef
|
||||
|
||||
gotNil := hasNoMemoryTool([]string{"no_mem_tool"}, nil)
|
||||
if gotNil {
|
||||
t.Error("hasNoMemoryTool with nil lookup should return false")
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
toolsUsed []string
|
||||
want bool
|
||||
}{
|
||||
{"empty tools", nil, false},
|
||||
{"no matching tool", []string{"unknown"}, false},
|
||||
{"tool without NoMemory", []string{"mem_tool"}, false},
|
||||
{"tool with NoMemory", []string{"no_mem_tool"}, true},
|
||||
{"mixed tools, first is no_memory", []string{"no_mem_tool", "mem_tool"}, true},
|
||||
{"mixed tools, last is no_memory", []string{"mem_tool", "no_mem_tool"}, true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := hasNoMemoryTool(tt.toolsUsed, lookup)
|
||||
if got != tt.want {
|
||||
t.Errorf("hasNoMemoryTool(%v) = %v, want %v", tt.toolsUsed, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func splitLines(s string) []string {
|
||||
var lines []string
|
||||
start := 0
|
||||
|
||||
@ -377,14 +377,15 @@ func docToTriples(doc *document.Doc) []memory.Triple {
|
||||
return triples
|
||||
}
|
||||
|
||||
func (a *Agent) emitMemoryCandidate(source, input, response string, toolsUsed []string) {
|
||||
func (a *Agent) emitMemoryCandidate(source, input, response string, toolResults []ToolResultItem, toolsUsed []string) {
|
||||
a.io.EmitOutput("memory", "memory_candidate", map[string]interface{}{
|
||||
"source": source,
|
||||
"input": input,
|
||||
"response": response,
|
||||
"tools_used": toolsUsed,
|
||||
"agent_id": string(a.id),
|
||||
"timestamp": time.Now().Unix(),
|
||||
"source": source,
|
||||
"input": input,
|
||||
"response": response,
|
||||
"tool_results": toolResults,
|
||||
"tools_used": toolsUsed,
|
||||
"agent_id": string(a.id),
|
||||
"timestamp": time.Now().Unix(),
|
||||
})
|
||||
}
|
||||
|
||||
@ -406,17 +407,18 @@ func (a *Agent) processConsolidation(evt *agentIO.InputEvent, input string) {
|
||||
Source: "system",
|
||||
Input: input,
|
||||
})
|
||||
response, toolsUsed, err := a.process(input, stageCtx)
|
||||
response, toolsUsed, toolResults, err := a.process(input, stageCtx)
|
||||
if err != nil {
|
||||
log.Printf("[agent] consolidation error: %v", err)
|
||||
return
|
||||
}
|
||||
a.context.Append(ContextEvent{
|
||||
Timestamp: time.Now(),
|
||||
Source: "agent",
|
||||
Input: input,
|
||||
Response: response,
|
||||
ToolsUsed: toolsUsed,
|
||||
Timestamp: time.Now(),
|
||||
Source: "agent",
|
||||
Input: input,
|
||||
Response: response,
|
||||
ToolsUsed: toolsUsed,
|
||||
ToolResults: toolResults,
|
||||
})
|
||||
log.Printf("[agent] consolidation done (%dms, tools=%v)", time.Since(start).Milliseconds(), toolsUsed)
|
||||
}
|
||||
|
||||
@ -183,7 +183,7 @@ func (a *Agent) processMediaInput(evt *agentIO.InputEvent) {
|
||||
Input: fallback,
|
||||
})
|
||||
|
||||
response, toolsUsed, err := a.process(fallback, stageCtx)
|
||||
response, toolsUsed, toolResults, err := a.process(fallback, stageCtx)
|
||||
if err != nil {
|
||||
log.Printf("[agent] process media error: %v", err)
|
||||
resp := fmt.Sprintf("处理错误: %v", err)
|
||||
@ -196,17 +196,18 @@ func (a *Agent) processMediaInput(evt *agentIO.InputEvent) {
|
||||
log.Printf("[agent] %s from %s → response (%dms, tools=%v)", evt.Type, evt.Source, elapsed.Milliseconds(), toolsUsed)
|
||||
|
||||
a.context.Append(ContextEvent{
|
||||
Timestamp: time.Now(),
|
||||
Source: "agent",
|
||||
Input: fallback,
|
||||
Response: response,
|
||||
ToolsUsed: toolsUsed,
|
||||
Timestamp: time.Now(),
|
||||
Source: "agent",
|
||||
Input: fallback,
|
||||
Response: response,
|
||||
ToolsUsed: toolsUsed,
|
||||
ToolResults: toolResults,
|
||||
})
|
||||
|
||||
a.emitResponse(evt, response)
|
||||
|
||||
if !stageCtx.NoMemory && !a.hasNoMemoryTool(toolsUsed) {
|
||||
a.emitMemoryCandidate(evt.Source, fallback, response, toolsUsed)
|
||||
if !stageCtx.NoMemory {
|
||||
a.emitMemoryCandidate(evt.Source, fallback, response, toolResults, toolsUsed)
|
||||
}
|
||||
}
|
||||
|
||||
@ -312,7 +313,7 @@ func (a *Agent) processTextInput(evt *agentIO.InputEvent, input string) {
|
||||
Input: input,
|
||||
})
|
||||
|
||||
response, toolsUsed, err := a.process(input, stageCtx)
|
||||
response, toolsUsed, toolResults, err := a.process(input, stageCtx)
|
||||
if err != nil {
|
||||
log.Printf("[agent] process error: %v", err)
|
||||
resp := fmt.Sprintf("处理错误: %v", err)
|
||||
@ -325,17 +326,18 @@ func (a *Agent) processTextInput(evt *agentIO.InputEvent, input string) {
|
||||
log.Printf("[agent] input from %s → response (%dms, tools=%v)", evt.Source, elapsed.Milliseconds(), toolsUsed)
|
||||
|
||||
a.context.Append(ContextEvent{
|
||||
Timestamp: time.Now(),
|
||||
Source: "agent",
|
||||
Input: input,
|
||||
Response: response,
|
||||
ToolsUsed: toolsUsed,
|
||||
Timestamp: time.Now(),
|
||||
Source: "agent",
|
||||
Input: input,
|
||||
Response: response,
|
||||
ToolsUsed: toolsUsed,
|
||||
ToolResults: toolResults,
|
||||
})
|
||||
|
||||
a.emitResponse(evt, response)
|
||||
|
||||
if !stageCtx.NoMemory && !a.hasNoMemoryTool(toolsUsed) {
|
||||
a.emitMemoryCandidate(evt.Source, input, response, toolsUsed)
|
||||
if !stageCtx.NoMemory {
|
||||
a.emitMemoryCandidate(evt.Source, input, response, toolResults, toolsUsed)
|
||||
}
|
||||
}
|
||||
|
||||
@ -386,15 +388,6 @@ func (a *Agent) emitResponse(evt *agentIO.InputEvent, response string) {
|
||||
a.runStage(sdk.StageAfterOutput, stageCtx)
|
||||
}
|
||||
|
||||
func (a *Agent) hasNoMemoryTool(toolsUsed []string) bool {
|
||||
for _, name := range toolsUsed {
|
||||
if def := a.stageHost.ToolDef(name); def != nil && def.NoMemory {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (a *Agent) drainInterrupts() []string {
|
||||
var out []string
|
||||
for {
|
||||
|
||||
@ -15,12 +15,12 @@ import (
|
||||
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
|
||||
)
|
||||
|
||||
func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response string, toolsUsed []string, err error) {
|
||||
func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response string, toolsUsed []string, toolResults []ToolResultItem, err error) {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
if a.provider == nil {
|
||||
return "", nil, fmt.Errorf("agent: no LLM provider configured")
|
||||
return "", nil, nil, fmt.Errorf("agent: no LLM provider configured")
|
||||
}
|
||||
|
||||
memContext := a.buildMemoryContext(input)
|
||||
@ -40,7 +40,7 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri
|
||||
a.docStoreSize())
|
||||
|
||||
if a.runStage(sdk.StagePreAction, stageCtx) {
|
||||
return *stageCtx.Response, toolsUsed, nil
|
||||
return *stageCtx.Response, toolsUsed, toolResults, nil
|
||||
}
|
||||
if len(stageCtx.ContextMsgs) > 0 {
|
||||
for _, m := range stageCtx.ContextMsgs {
|
||||
@ -132,11 +132,11 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri
|
||||
if llmErr != nil {
|
||||
if errors.Is(llmErr, context.Canceled) && a.ctx.Err() == nil {
|
||||
if a.currentOutputChannel == "_consolidation_" {
|
||||
return "", toolsUsed, fmt.Errorf("interrupted by user input")
|
||||
return "", toolsUsed, toolResults, fmt.Errorf("interrupted by user input")
|
||||
}
|
||||
continue
|
||||
}
|
||||
return "", toolsUsed, fmt.Errorf("all %d providers failed, last error: %w",
|
||||
return "", toolsUsed, toolResults, fmt.Errorf("all %d providers failed, last error: %w",
|
||||
len(providers), llmErr)
|
||||
}
|
||||
|
||||
@ -154,7 +154,7 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri
|
||||
}
|
||||
}
|
||||
if a.runStage(sdk.StagePostAction, stageCtx) {
|
||||
return *stageCtx.Response, toolsUsed, nil
|
||||
return *stageCtx.Response, toolsUsed, toolResults, nil
|
||||
}
|
||||
resp.Content = stageCtx.LLMText
|
||||
resp.ToolCalls = convertBackToolCalls(stageCtx.ToolCalls)
|
||||
@ -176,7 +176,7 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri
|
||||
a.publishEvent(events.EventAgentLLMChain, chainPayload)
|
||||
|
||||
if len(resp.ToolCalls) == 0 {
|
||||
return resp.Content, toolsUsed, nil
|
||||
return resp.Content, toolsUsed, toolResults, nil
|
||||
}
|
||||
|
||||
contentOnce := true
|
||||
@ -226,6 +226,7 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri
|
||||
}
|
||||
|
||||
result := a.executeToolCall(tc)
|
||||
toolResults = append(toolResults, ToolResultItem{Name: tc.Name, Output: result})
|
||||
log.Printf("[agent] tool %s result: %s", tc.Name, truncateStr(result, 100))
|
||||
|
||||
stageCtx.ToolResults = []sdk.ToolResult{{CallID: tc.ID, Name: tc.Name, Plugin: pluginName, Success: true, Result: result}}
|
||||
|
||||
@ -159,6 +159,29 @@ func (h *StageHost) RunStage(stage sdk.Stage, ctx *sdk.StageContext) {
|
||||
}
|
||||
}
|
||||
|
||||
func (h *StageHost) ToolDefCleaner(name string) func(string) string {
|
||||
h.mu.RLock()
|
||||
defer h.mu.RUnlock()
|
||||
for _, def := range h.toolDefs {
|
||||
if def.Name == name {
|
||||
return def.Cleaner
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *StageHost) NoMemoryToolNames() map[string]bool {
|
||||
h.mu.RLock()
|
||||
defer h.mu.RUnlock()
|
||||
set := make(map[string]bool, len(h.toolDefs))
|
||||
for _, def := range h.toolDefs {
|
||||
if def.NoMemory {
|
||||
set[def.Name] = true
|
||||
}
|
||||
}
|
||||
return set
|
||||
}
|
||||
|
||||
func (h *StageHost) ToolCount() int {
|
||||
h.mu.RLock()
|
||||
defer h.mu.RUnlock()
|
||||
|
||||
@ -7,24 +7,24 @@ import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func init() {
|
||||
globalTextCleaner = func(text string) string {
|
||||
reQQGroupSuffix := regexp.MustCompile(`,通过id\d+使用qq_get_message工具获取消息正文。获取内容后使用 output_send\(channel="qq"\) 回复该群聊,content 设为 JSON 字符串:\{[^}]*\}`)
|
||||
reQQPrivateSuffix := regexp.MustCompile(`,通过id\d+使用qq_get_message工具获取消息正文。获取内容后使用 output_send\(channel="qq"\) 回复对方,content 设为 JSON 字符串:\{[^}]*\}`)
|
||||
reQQOldReply := regexp.MustCompile(`通过id\d+使用qq_get_message工具获取消息正文。获取后必须使用[^。]+。`)
|
||||
reQQOldForbid := regexp.MustCompile(`你只能通过qq_get_message先看消息,然后直接用%!s\(MISSING\)send_private_msg回复,中间的思考过程禁止调用任何其他工具\s*→\s*`)
|
||||
reQQGeneral := regexp.MustCompile(`通过id\d+使用qq_get_message工具获取消息正文[。,][^。]*?(?:回复|发送消息)`)
|
||||
reTimestamp := regexp.MustCompile(`\[\d{2}:\d{2}\]\s*`)
|
||||
reMultiSpace := regexp.MustCompile(`\s+`)
|
||||
text = reQQGroupSuffix.ReplaceAllString(text, "")
|
||||
text = reQQPrivateSuffix.ReplaceAllString(text, "")
|
||||
text = reQQOldReply.ReplaceAllString(text, "")
|
||||
text = reQQOldForbid.ReplaceAllString(text, "")
|
||||
text = reQQGeneral.ReplaceAllString(text, "")
|
||||
text = reTimestamp.ReplaceAllString(text, "")
|
||||
text = reMultiSpace.ReplaceAllString(text, " ")
|
||||
return text
|
||||
}
|
||||
// cleanQQTemplate 模拟之前由 globalTextCleaner 执行的模板噪音清理,
|
||||
// 用于 stress test 中生成 cleanedText。
|
||||
func cleanQQTemplate(text string) string {
|
||||
reQQGroupSuffix := regexp.MustCompile(`,通过id\d+使用qq_get_message工具获取消息正文。获取内容后使用 output_send\(channel="qq"\) 回复该群聊,content 设为 JSON 字符串:\{[^}]*\}`)
|
||||
reQQPrivateSuffix := regexp.MustCompile(`,通过id\d+使用qq_get_message工具获取消息正文。获取内容后使用 output_send\(channel="qq"\) 回复对方,content 设为 JSON 字符串:\{[^}]*\}`)
|
||||
reQQOldReply := regexp.MustCompile(`通过id\d+使用qq_get_message工具获取消息正文。获取后必须使用[^。]+。`)
|
||||
reQQOldForbid := regexp.MustCompile(`你只能通过qq_get_message先看消息,然后直接用%!s\(MISSING\)send_private_msg回复对方,中间的思考过程禁止调用任何其他工具\s*→\s*`)
|
||||
reQQGeneral := regexp.MustCompile(`通过id\d+使用qq_get_message工具获取消息正文[。,][^。]*?(?:回复|发送消息)`)
|
||||
reTimestamp := regexp.MustCompile(`\[\d{2}:\d{2}\]\s*`)
|
||||
reMultiSpace := regexp.MustCompile(`\s+`)
|
||||
text = reQQGroupSuffix.ReplaceAllString(text, "")
|
||||
text = reQQPrivateSuffix.ReplaceAllString(text, "")
|
||||
text = reQQOldReply.ReplaceAllString(text, "")
|
||||
text = reQQOldForbid.ReplaceAllString(text, "")
|
||||
text = reQQGeneral.ReplaceAllString(text, "")
|
||||
text = reTimestamp.ReplaceAllString(text, "")
|
||||
text = reMultiSpace.ReplaceAllString(text, " ")
|
||||
return text
|
||||
}
|
||||
|
||||
type cleanTestEvent struct {
|
||||
@ -305,13 +305,14 @@ func genStressEvents(n int) []cleanTestEvent {
|
||||
}
|
||||
|
||||
func cleanEventText(source, input, response string) string {
|
||||
// 先做基础 CleanText(去空格/逗号),再做模板噪音清理
|
||||
switch {
|
||||
case source == "agent" && response != "":
|
||||
return CleanText(response)
|
||||
return cleanQQTemplate(CleanText(response))
|
||||
case source == "cold_storage":
|
||||
return CleanText(input + " " + response)
|
||||
return cleanQQTemplate(CleanText(input + " " + response))
|
||||
default:
|
||||
return CleanText(input)
|
||||
return cleanQQTemplate(CleanText(input))
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@ -6,17 +6,7 @@ import (
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
|
||||
)
|
||||
|
||||
var globalTextCleaner func(string) string
|
||||
|
||||
func SetTextCleaner(fn func(string) string) {
|
||||
globalTextCleaner = fn
|
||||
}
|
||||
|
||||
func CleanText(text string) string {
|
||||
if globalTextCleaner != nil {
|
||||
text = globalTextCleaner(text)
|
||||
}
|
||||
|
||||
text = strings.TrimSpace(text)
|
||||
|
||||
if text == "" {
|
||||
|
||||
@ -5,10 +5,6 @@ import (
|
||||
)
|
||||
|
||||
func TestCleanTextTrim(t *testing.T) {
|
||||
prev := globalTextCleaner
|
||||
globalTextCleaner = nil
|
||||
defer func() { globalTextCleaner = prev }()
|
||||
|
||||
tests := []struct {
|
||||
input string
|
||||
expected string
|
||||
@ -29,55 +25,3 @@ func TestCleanTextTrim(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanTextWithRegisteredCleaner(t *testing.T) {
|
||||
prev := globalTextCleaner
|
||||
globalTextCleaner = func(text string) string {
|
||||
return "prefix_" + text
|
||||
}
|
||||
defer func() { globalTextCleaner = prev }()
|
||||
|
||||
got := CleanText(" hello ")
|
||||
if got != "prefix_ hello" {
|
||||
t.Errorf("CleanText with cleaner = %q, want %q", got, "prefix_ hello")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanTextCleanerChain(t *testing.T) {
|
||||
prev := globalTextCleaner
|
||||
globalTextCleaner = func(text string) string {
|
||||
text = text + "_step1"
|
||||
text = text + "_step2"
|
||||
return text
|
||||
}
|
||||
defer func() { globalTextCleaner = prev }()
|
||||
|
||||
got := CleanText("test")
|
||||
if got != "test_step1_step2" {
|
||||
t.Errorf("CleanText chain = %q, want %q", got, "test_step1_step2")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetTextCleanerReplace(t *testing.T) {
|
||||
prev := globalTextCleaner
|
||||
globalTextCleaner = func(text string) string { return "old_" + text }
|
||||
|
||||
SetTextCleaner(func(text string) string { return "new_" + text })
|
||||
defer func() { globalTextCleaner = prev }()
|
||||
|
||||
got := CleanText("x")
|
||||
if got != "new_x" {
|
||||
t.Errorf("after SetTextCleaner = %q, want %q", got, "new_x")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanTextEmptyAfterCleaner(t *testing.T) {
|
||||
prev := globalTextCleaner
|
||||
globalTextCleaner = func(text string) string { return "" }
|
||||
defer func() { globalTextCleaner = prev }()
|
||||
|
||||
got := CleanText("something")
|
||||
if got != "" {
|
||||
t.Errorf("expected empty, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
@ -124,25 +124,34 @@ func (s *Store) Insert(doc *Doc) error {
|
||||
}
|
||||
|
||||
// ContextToDoc — 将一段上下文对话历史提炼为文档(带内容去重)
|
||||
func (s *Store) ContextToDoc(source string, entries []ContextEntry, vec vector.Vectorizer) (*Doc, error) {
|
||||
// cleanFn 可选,用于在计算层(摘要/标签/实体提取)前过滤文本,不影响原文存储。
|
||||
func (s *Store) ContextToDoc(source string, entries []ContextEntry, vec vector.Vectorizer, cleanFn ...func(string) string) (*Doc, error) {
|
||||
if len(entries) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
cleanText := func(text string) string { return text }
|
||||
if len(cleanFn) > 0 && cleanFn[0] != nil {
|
||||
cleanText = cleanFn[0]
|
||||
}
|
||||
|
||||
var parts []string
|
||||
for _, e := range entries {
|
||||
line := fmt.Sprintf("[%s] %s: %s", e.Timestamp.Format("15:04"), e.Source, e.Content)
|
||||
if e.Response != "" {
|
||||
line += fmt.Sprintf(" → %s", truncate(e.Response, 100))
|
||||
}
|
||||
for _, tr := range e.ToolResults {
|
||||
line += fmt.Sprintf("\n [工具] %s: %s", tr.Name, truncate(tr.Output, 200))
|
||||
}
|
||||
parts = append(parts, line)
|
||||
}
|
||||
content := strings.Join(parts, "\n")
|
||||
contentHash := simpleHash(content)
|
||||
|
||||
summary := summarizeEntries(entries)
|
||||
tags := extractTags(entries)
|
||||
entities := extractEntities(entries)
|
||||
summary := summarizeEntries(entries, cleanText)
|
||||
tags := extractTags(entries, cleanText)
|
||||
entities := extractEntities(entries, cleanText)
|
||||
|
||||
s.mu.Lock()
|
||||
|
||||
@ -417,23 +426,38 @@ func (s *Store) flush() {
|
||||
s.dirty = false
|
||||
}
|
||||
|
||||
type ContextEntry struct {
|
||||
Timestamp time.Time
|
||||
Source string
|
||||
Content string
|
||||
Response string
|
||||
type ToolResultItem struct {
|
||||
Name string
|
||||
Output string
|
||||
}
|
||||
|
||||
func summarizeEntries(entries []ContextEntry) string {
|
||||
type ContextEntry struct {
|
||||
Timestamp time.Time
|
||||
Source string
|
||||
Content string
|
||||
Response string
|
||||
ToolResults []ToolResultItem
|
||||
}
|
||||
|
||||
func summarizeEntries(entries []ContextEntry, cleanText ...func(string) string) string {
|
||||
if len(entries) == 0 {
|
||||
return ""
|
||||
}
|
||||
clean := func(text string) string { return text }
|
||||
if len(cleanText) > 0 && cleanText[0] != nil {
|
||||
clean = cleanText[0]
|
||||
}
|
||||
sources := make(map[string]int)
|
||||
var topics []string
|
||||
for _, e := range entries {
|
||||
sources[e.Source]++
|
||||
words := memory.ExtractKeywords(e.Content)
|
||||
words := memory.ExtractKeywords(clean(e.Content))
|
||||
topics = append(topics, words...)
|
||||
for _, tr := range e.ToolResults {
|
||||
cleaned := clean(tr.Output)
|
||||
toolWords := memory.ExtractKeywords(cleaned)
|
||||
topics = append(topics, toolWords...)
|
||||
}
|
||||
}
|
||||
|
||||
summary := fmt.Sprintf("来自 %d 个来源的 %d 条对话", len(sources), len(entries))
|
||||
@ -461,12 +485,21 @@ func summarizeEntries(entries []ContextEntry) string {
|
||||
return summary
|
||||
}
|
||||
|
||||
func extractTags(entries []ContextEntry) []string {
|
||||
func extractTags(entries []ContextEntry, cleanText ...func(string) string) []string {
|
||||
clean := func(text string) string { return text }
|
||||
if len(cleanText) > 0 && cleanText[0] != nil {
|
||||
clean = cleanText[0]
|
||||
}
|
||||
tagSet := make(map[string]bool)
|
||||
for _, e := range entries {
|
||||
for _, kw := range memory.ExtractKeywords(e.Content) {
|
||||
for _, kw := range memory.ExtractKeywords(clean(e.Content)) {
|
||||
tagSet[kw] = true
|
||||
}
|
||||
for _, tr := range e.ToolResults {
|
||||
for _, kw := range memory.ExtractKeywords(clean(tr.Output)) {
|
||||
tagSet[kw] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
var tags []string
|
||||
for t := range tagSet {
|
||||
@ -478,17 +511,29 @@ func extractTags(entries []ContextEntry) []string {
|
||||
return tags
|
||||
}
|
||||
|
||||
func extractEntities(entries []ContextEntry) []string {
|
||||
func extractEntities(entries []ContextEntry, cleanText ...func(string) string) []string {
|
||||
// 简易实体提取:提取引号内的内容、粗体/标记词
|
||||
clean := func(text string) string { return text }
|
||||
if len(cleanText) > 0 && cleanText[0] != nil {
|
||||
clean = cleanText[0]
|
||||
}
|
||||
var entities []string
|
||||
seen := make(map[string]bool)
|
||||
for _, e := range entries {
|
||||
for _, kw := range memory.ExtractKeywords(e.Content) {
|
||||
for _, kw := range memory.ExtractKeywords(clean(e.Content)) {
|
||||
if len(kw) >= 2 && !seen[kw] {
|
||||
seen[kw] = true
|
||||
entities = append(entities, kw)
|
||||
}
|
||||
}
|
||||
for _, tr := range e.ToolResults {
|
||||
for _, kw := range memory.ExtractKeywords(clean(tr.Output)) {
|
||||
if len(kw) >= 2 && !seen[kw] {
|
||||
seen[kw] = true
|
||||
entities = append(entities, kw)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(entities) > 20 {
|
||||
entities = entities[:20]
|
||||
|
||||
@ -53,7 +53,7 @@ func TestCleanText(t *testing.T) {
|
||||
}
|
||||
|
||||
for i, c := range cases {
|
||||
got := CleanText(c.input)
|
||||
got := cleanQQTemplate(CleanText(c.input))
|
||||
if c.expected != "" && got != c.expected {
|
||||
t.Errorf("case %d:\n input: %q\n expected: %q\n got: %q", i, trimLen(c.input, 60), c.expected, got)
|
||||
}
|
||||
|
||||
@ -84,7 +84,6 @@ type Registry struct {
|
||||
toolCleaner PluginToolCleaner
|
||||
|
||||
knownDisabled map[string]bool
|
||||
textCleaners []func(string) string
|
||||
}
|
||||
|
||||
func NewRegistry() *Registry {
|
||||
@ -110,17 +109,6 @@ func (r *Registry) SetStageRegistrar(fn sdk.StageRegistrar) { r.regStage =
|
||||
func (r *Registry) SetAPIRegistrar(fn sdk.APIRegistrar) { r.regAPI = fn }
|
||||
func (r *Registry) SetToolCleaner(tc PluginToolCleaner) { r.toolCleaner = tc }
|
||||
|
||||
// CleanText applies all registered text cleaners in order.
|
||||
func (r *Registry) CleanText(text string) string {
|
||||
r.mu.RLock()
|
||||
cleaners := r.textCleaners
|
||||
r.mu.RUnlock()
|
||||
for _, fn := range cleaners {
|
||||
text = fn(text)
|
||||
}
|
||||
return text
|
||||
}
|
||||
|
||||
func (r *Registry) RegisterNative(name string, factory NativeFactory) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
@ -270,7 +258,6 @@ func (r *Registry) Load(dir string) error {
|
||||
r.plugins[name] = p
|
||||
r.pluginAutoRestart[name] = plgSDK.AutoRestart()
|
||||
r.instances = append(r.instances, p)
|
||||
r.textCleaners = append(r.textCleaners, plgSDK.TextCleaners()...)
|
||||
r.mu.Unlock()
|
||||
log.Printf("[plugin] loaded: %s", name)
|
||||
}
|
||||
@ -353,7 +340,6 @@ func (r *Registry) loadOne(plgDir, name string) bool {
|
||||
r.plugins[name] = plg
|
||||
r.pluginAutoRestart[name] = plgSDK.AutoRestart()
|
||||
r.instances = append(r.instances, plg)
|
||||
r.textCleaners = append(r.textCleaners, plgSDK.TextCleaners()...)
|
||||
r.mu.Unlock()
|
||||
log.Printf("[plugin] loaded: %s", name)
|
||||
return true
|
||||
@ -370,7 +356,6 @@ func (r *Registry) StopAll() {
|
||||
r.plugins = make(map[string]sdk.Plugin)
|
||||
r.instances = nil
|
||||
r.pluginAutoRestart = make(map[string]bool)
|
||||
r.textCleaners = nil
|
||||
}
|
||||
|
||||
func (r *Registry) Reload(dir string) (string, error) {
|
||||
|
||||
@ -214,8 +214,13 @@ func TestCmdRunNonZeroExit(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestTruncateOutput(t *testing.T) {
|
||||
p, _, err := setupPlugin()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
short := "hello"
|
||||
if s := truncateOutput(short); s != short {
|
||||
if s := p.truncateOutput(short); s != short {
|
||||
t.Fatalf("expected %q, got %q", short, s)
|
||||
}
|
||||
|
||||
@ -223,7 +228,7 @@ func TestTruncateOutput(t *testing.T) {
|
||||
for i := range long {
|
||||
long[i] = 'x'
|
||||
}
|
||||
s := truncateOutput(string(long))
|
||||
s := p.truncateOutput(string(long))
|
||||
if len(s) >= 40000 {
|
||||
t.Fatal("expected truncation")
|
||||
}
|
||||
|
||||
@ -1,6 +1,7 @@
|
||||
package files
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
@ -61,6 +62,14 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
|
||||
s.RegisterTool(tp+"read", sdk.ToolDef{
|
||||
Name: tp + "read",
|
||||
Description: fmt.Sprintf("读取文件内容。支持 offset/limit 分段读取大文件。沙箱路径: %s", p.filesDir),
|
||||
NoMemory: false,
|
||||
Cleaner: func(output string) string {
|
||||
var r struct{ Content string }
|
||||
if err := json.Unmarshal([]byte(output), &r); err != nil {
|
||||
return output
|
||||
}
|
||||
return r.Content
|
||||
},
|
||||
Parameters: map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{
|
||||
|
||||
Reference in New Issue
Block a user