mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-03 15:53:56 +00:00
refactor: pluginize text cleaning and tool NoMemory control
- SDK: ToolDef.NoMemory field, PluginSDK.RegisterTextCleaner/TextCleaners - Registry: aggregate text cleaners from plugins, expose CleanText() - Memory: replace hardcoded QQ regex CleanTemplateText with dynamic CleanText/SetTextCleaner - StageHost: add ToolDef(name) lookup - eventloop: check ToolDef.NoMemory before emitMemoryCandidate - context/Prune: replace hardcoded agentcli/terminal source filter with ToolsUsed NoMemory check - agentcli/cmd: mark tools with NoMemory: true - main.go: wire memory.SetTextCleaner(pluginReg.CleanText)
This commit is contained in:
@ -165,6 +165,11 @@ func New(cfg AgentConfig) *Agent {
|
||||
cfg.DocStore.ReindexWithVectorizer(embedder)
|
||||
}
|
||||
|
||||
rc := NewRelevanceContext(cfg.ContextSavePath, embedder)
|
||||
if cfg.StageHost != nil {
|
||||
rc.SetToolDefLookup(cfg.StageHost.ToolDef)
|
||||
}
|
||||
|
||||
return &Agent{
|
||||
id: cfg.ID,
|
||||
startTime: time.Now(),
|
||||
@ -175,7 +180,7 @@ func New(cfg AgentConfig) *Agent {
|
||||
indexer: cfg.Indexer,
|
||||
skills: cfg.Skills,
|
||||
tracker: cfg.Tracker,
|
||||
context: NewRelevanceContext(cfg.ContextSavePath, embedder),
|
||||
context: rc,
|
||||
systemPrompt: cfg.SystemPrompt,
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
|
||||
@ -13,6 +13,7 @@ import (
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/document"
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
|
||||
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
|
||||
)
|
||||
|
||||
type ContextEvent struct {
|
||||
@ -27,12 +28,13 @@ type ContextEvent struct {
|
||||
const contextFlushInterval = 5 * time.Second
|
||||
|
||||
type RelevanceContext struct {
|
||||
mu sync.Mutex
|
||||
events []*ContextEvent
|
||||
embedder *memory.StaticEmbedder
|
||||
savePath string
|
||||
saveTimer *time.Timer
|
||||
dirty bool
|
||||
mu sync.Mutex
|
||||
events []*ContextEvent
|
||||
embedder *memory.StaticEmbedder
|
||||
savePath string
|
||||
saveTimer *time.Timer
|
||||
dirty bool
|
||||
toolDefLookup func(name string) *sdk.ToolDef
|
||||
}
|
||||
|
||||
func NewRelevanceContext(savePath string, embedder *memory.StaticEmbedder) *RelevanceContext {
|
||||
@ -46,6 +48,12 @@ func NewRelevanceContext(savePath string, embedder *memory.StaticEmbedder) *Rele
|
||||
return rc
|
||||
}
|
||||
|
||||
func (c *RelevanceContext) SetToolDefLookup(fn func(name string) *sdk.ToolDef) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.toolDefLookup = fn
|
||||
}
|
||||
|
||||
func (c *RelevanceContext) load() {
|
||||
data, err := os.ReadFile(c.savePath)
|
||||
if err != nil {
|
||||
@ -56,7 +64,7 @@ func (c *RelevanceContext) load() {
|
||||
return
|
||||
}
|
||||
for _, evt := range events {
|
||||
evt.Input = memory.CleanTemplateText(evt.Input)
|
||||
evt.Input = memory.CleanText(evt.Input)
|
||||
evt.Vector = c.computeVector(evt)
|
||||
}
|
||||
c.events = events
|
||||
@ -65,11 +73,11 @@ func (c *RelevanceContext) load() {
|
||||
func textForVector(evt *ContextEvent) string {
|
||||
switch {
|
||||
case evt.Source == "agent" && evt.Response != "":
|
||||
return memory.CleanTemplateText(evt.Response)
|
||||
return memory.CleanText(evt.Response)
|
||||
case evt.Source == "cold_storage":
|
||||
return memory.CleanTemplateText(evt.Input + " " + evt.Response)
|
||||
return memory.CleanText(evt.Input + " " + evt.Response)
|
||||
default:
|
||||
return memory.CleanTemplateText(evt.Input)
|
||||
return memory.CleanText(evt.Input)
|
||||
}
|
||||
}
|
||||
|
||||
@ -95,7 +103,7 @@ func (c *RelevanceContext) Append(evt ContextEvent) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
evt.Input = memory.CleanTemplateText(evt.Input)
|
||||
evt.Input = memory.CleanText(evt.Input)
|
||||
evt.Vector = c.computeVector(&evt)
|
||||
c.events = append(c.events, &evt)
|
||||
|
||||
@ -193,7 +201,7 @@ func (c *RelevanceContext) Prune(currentInput string, topK int, docStore *docume
|
||||
if docStore != nil && len(archive) > 0 {
|
||||
var filtered []scored
|
||||
for _, s := range archive {
|
||||
if s.event.Source == "agentcli" || s.event.Source == "terminal" {
|
||||
if hasNoMemoryTool(s.event.ToolsUsed, c.toolDefLookup) {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, s)
|
||||
@ -257,3 +265,15 @@ func (c *RelevanceContext) Len() int {
|
||||
defer c.mu.Unlock()
|
||||
return len(c.events)
|
||||
}
|
||||
|
||||
func hasNoMemoryTool(toolsUsed []string, lookup func(string) *sdk.ToolDef) bool {
|
||||
if lookup == nil {
|
||||
return false
|
||||
}
|
||||
for _, name := range toolsUsed {
|
||||
if def := lookup(name); def != nil && def.NoMemory {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@ -6,6 +6,7 @@ import (
|
||||
"time"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
|
||||
)
|
||||
|
||||
func newTestCtx() *RelevanceContext {
|
||||
@ -237,6 +238,40 @@ 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
|
||||
|
||||
@ -205,7 +205,7 @@ func (a *Agent) processMediaInput(evt *agentIO.InputEvent) {
|
||||
|
||||
a.emitResponse(evt, response)
|
||||
|
||||
if !stageCtx.NoMemory {
|
||||
if !stageCtx.NoMemory && !a.hasNoMemoryTool(toolsUsed) {
|
||||
a.emitMemoryCandidate(evt.Source, fallback, response, toolsUsed)
|
||||
}
|
||||
}
|
||||
@ -334,7 +334,7 @@ func (a *Agent) processTextInput(evt *agentIO.InputEvent, input string) {
|
||||
|
||||
a.emitResponse(evt, response)
|
||||
|
||||
if !stageCtx.NoMemory {
|
||||
if !stageCtx.NoMemory && !a.hasNoMemoryTool(toolsUsed) {
|
||||
a.emitMemoryCandidate(evt.Source, input, response, toolsUsed)
|
||||
}
|
||||
}
|
||||
@ -386,6 +386,15 @@ 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 {
|
||||
|
||||
@ -54,6 +54,17 @@ func (h *StageHost) GetToolDefs() []sdk.ToolDef {
|
||||
return defs
|
||||
}
|
||||
|
||||
func (h *StageHost) ToolDef(name string) *sdk.ToolDef {
|
||||
h.mu.RLock()
|
||||
defer h.mu.RUnlock()
|
||||
for _, def := range h.toolDefs {
|
||||
if def.Name == name {
|
||||
return &def
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *StageHost) ExecuteTool(name string, args map[string]interface{}) (ret interface{}, err error) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
|
||||
@ -197,3 +197,54 @@ func TestStageHostMultipleTools(t *testing.T) {
|
||||
t.Errorf("expected from_p2, got %v", r2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStageHostToolDefLookup(t *testing.T) {
|
||||
host := NewStageHost()
|
||||
host.RegisterTool("tool_a", sdk.ToolDef{Name: "tool_a", NoMemory: true}, nil)
|
||||
host.RegisterTool("tool_b", sdk.ToolDef{Name: "tool_b"}, nil)
|
||||
|
||||
def := host.ToolDef("tool_a")
|
||||
if def == nil {
|
||||
t.Fatal("expected tool_a to be found")
|
||||
}
|
||||
if !def.NoMemory {
|
||||
t.Error("tool_a should have NoMemory=true")
|
||||
}
|
||||
|
||||
def = host.ToolDef("tool_b")
|
||||
if def == nil {
|
||||
t.Fatal("expected tool_b to be found")
|
||||
}
|
||||
if def.NoMemory {
|
||||
t.Error("tool_b should have NoMemory=false")
|
||||
}
|
||||
|
||||
def = host.ToolDef("nonexistent")
|
||||
if def != nil {
|
||||
t.Errorf("expected nil for nonexistent tool, got %v", def)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStageHostToolDefNoMemoryStored(t *testing.T) {
|
||||
host := NewStageHost()
|
||||
host.RegisterTool("mem_tool", sdk.ToolDef{Name: "mem_tool", NoMemory: true}, nil)
|
||||
host.RegisterTool("normal_tool", sdk.ToolDef{Name: "normal_tool", NoMemory: false}, nil)
|
||||
|
||||
defs := host.GetToolDefs()
|
||||
found := map[string]bool{}
|
||||
for _, d := range defs {
|
||||
found[d.Name] = d.NoMemory
|
||||
}
|
||||
|
||||
if v, ok := found["mem_tool"]; !ok {
|
||||
t.Error("mem_tool not found in defs")
|
||||
} else if !v {
|
||||
t.Error("mem_tool.NoMemory should be true")
|
||||
}
|
||||
|
||||
if v, ok := found["normal_tool"]; !ok {
|
||||
t.Error("normal_tool not found in defs")
|
||||
} else if v {
|
||||
t.Error("normal_tool.NoMemory should be false")
|
||||
}
|
||||
}
|
||||
|
||||
@ -180,7 +180,7 @@ func TestBilingualVectorizeClean(t *testing.T) {
|
||||
cleanVec := e.VectorizeClean(inp)
|
||||
sim := cosineSim(rawVec, cleanVec)
|
||||
rawTokens := len(e.tokenize(inp))
|
||||
cleanTokens := len(e.tokenize(CleanTemplateText(inp)))
|
||||
cleanTokens := len(e.tokenize(CleanText(inp)))
|
||||
t.Logf("[%d] sim(raw,clean)=%.4f tokens: raw=%d clean=%d", i, sim, rawTokens, cleanTokens)
|
||||
}
|
||||
}
|
||||
@ -236,8 +236,8 @@ func genBilingualEvents() []bilingualEvent {
|
||||
func textForBilingual(ev bilingualEvent, modelPaths []string) string {
|
||||
switch {
|
||||
case ev.source == "agent" && ev.text != "":
|
||||
return CleanTemplateText(ev.text)
|
||||
return CleanText(ev.text)
|
||||
default:
|
||||
return CleanTemplateText(ev.text)
|
||||
return CleanText(ev.text)
|
||||
}
|
||||
}
|
||||
|
||||
@ -2,10 +2,31 @@ package memory
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"sort"
|
||||
"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
|
||||
}
|
||||
}
|
||||
|
||||
type cleanTestEvent struct {
|
||||
idx int
|
||||
source string
|
||||
@ -173,7 +194,7 @@ func TestCleanVectorConsistency(t *testing.T) {
|
||||
}
|
||||
|
||||
for _, tmpl := range templates {
|
||||
cleaned := CleanTemplateText(tmpl)
|
||||
cleaned := CleanText(tmpl)
|
||||
t.Logf("template {%q} → {%q} (%d chars)", trimLen(tmpl, 60), cleaned, len(cleaned))
|
||||
}
|
||||
|
||||
@ -286,11 +307,11 @@ func genStressEvents(n int) []cleanTestEvent {
|
||||
func cleanEventText(source, input, response string) string {
|
||||
switch {
|
||||
case source == "agent" && response != "":
|
||||
return CleanTemplateText(response)
|
||||
return CleanText(response)
|
||||
case source == "cold_storage":
|
||||
return CleanTemplateText(input + " " + response)
|
||||
return CleanText(input + " " + response)
|
||||
default:
|
||||
return CleanTemplateText(input)
|
||||
return CleanText(input)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@ -1,45 +1,22 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
|
||||
)
|
||||
|
||||
var (
|
||||
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*`,
|
||||
)
|
||||
reAgentPrefix = regexp.MustCompile(
|
||||
`冷知识|注意|提示|核心要求|规则`,
|
||||
)
|
||||
reMultiSpace = regexp.MustCompile(`\s+`)
|
||||
)
|
||||
var globalTextCleaner func(string) string
|
||||
|
||||
func SetTextCleaner(fn func(string) string) {
|
||||
globalTextCleaner = fn
|
||||
}
|
||||
|
||||
func CleanText(text string) string {
|
||||
if globalTextCleaner != nil {
|
||||
text = globalTextCleaner(text)
|
||||
}
|
||||
|
||||
func CleanTemplateText(text string) string {
|
||||
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, " ")
|
||||
text = strings.TrimSpace(text)
|
||||
|
||||
if text == "" {
|
||||
@ -53,5 +30,5 @@ func CleanTemplateText(text string) string {
|
||||
}
|
||||
|
||||
func (e *StaticEmbedder) VectorizeClean(text string) vector.Vector {
|
||||
return e.Vectorize(CleanTemplateText(text))
|
||||
return e.Vectorize(CleanText(text))
|
||||
}
|
||||
|
||||
83
internal/memory/clean_text_test.go
Normal file
83
internal/memory/clean_text_test.go
Normal file
@ -0,0 +1,83 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCleanTextTrim(t *testing.T) {
|
||||
prev := globalTextCleaner
|
||||
globalTextCleaner = nil
|
||||
defer func() { globalTextCleaner = prev }()
|
||||
|
||||
tests := []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{" hello ", "hello"},
|
||||
{",hello", "hello"},
|
||||
{",,hello", "hello"},
|
||||
{" ,,hello ", "hello"},
|
||||
{"", ""},
|
||||
{" ", ""},
|
||||
{",", ""},
|
||||
{",x", "x"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
got := CleanText(tt.input)
|
||||
if got != tt.expected {
|
||||
t.Errorf("CleanText(%q) = %q, want %q", tt.input, got, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
@ -93,7 +93,7 @@ var stopWords = map[string]bool{
|
||||
}
|
||||
|
||||
func ExtractKeywords(text string) []string {
|
||||
text = CleanTemplateText(text)
|
||||
text = CleanText(text)
|
||||
x := GetJieba()
|
||||
if x == nil {
|
||||
return nil
|
||||
@ -120,7 +120,7 @@ func ExtractKeywords(text string) []string {
|
||||
|
||||
// CutExact 精确模式分词:返回去停用词后的所有有义项(不限数量),用于 doc→graph 蒸馏
|
||||
func CutExact(text string) []string {
|
||||
text = CleanTemplateText(text)
|
||||
text = CleanText(text)
|
||||
x := GetJieba()
|
||||
if x == nil {
|
||||
return nil
|
||||
|
||||
@ -101,7 +101,7 @@ func TestCutExactRemoveTimestamp(t *testing.T) {
|
||||
got := CutExact("[15:04] 今天天气不错")
|
||||
for _, g := range got {
|
||||
if g == "15" || g == "04" || g == "15:04" {
|
||||
t.Errorf("timestamp should be removed by CleanTemplateText, got %q in %v", g, got)
|
||||
t.Errorf("timestamp should be removed by CleanText, got %q in %v", g, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@ -90,7 +90,7 @@ func (idx *Indexer) BuildContext(userInput string) *InjectedContext {
|
||||
return &InjectedContext{Summary: ""}
|
||||
}
|
||||
|
||||
input := CleanTemplateText(userInput)
|
||||
input := CleanText(userInput)
|
||||
|
||||
// 1. 向量搜索:从实体名向量索引中找到相关实体
|
||||
vectorEntities := idx.vectorSearchEntities(input)
|
||||
|
||||
@ -16,7 +16,7 @@ type realEvent struct {
|
||||
Response string `json:"response"`
|
||||
}
|
||||
|
||||
func TestCleanTemplateText(t *testing.T) {
|
||||
func TestCleanText(t *testing.T) {
|
||||
cases := []struct {
|
||||
input string
|
||||
expected string
|
||||
@ -53,7 +53,7 @@ func TestCleanTemplateText(t *testing.T) {
|
||||
}
|
||||
|
||||
for i, c := range cases {
|
||||
got := CleanTemplateText(c.input)
|
||||
got := 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)
|
||||
}
|
||||
@ -91,11 +91,11 @@ func TestRealContextPerSourceVector(t *testing.T) {
|
||||
clean := func(ev realEvent) string {
|
||||
switch {
|
||||
case ev.Source == "agent" && ev.Response != "":
|
||||
return CleanTemplateText(ev.Response)
|
||||
return CleanText(ev.Response)
|
||||
case ev.Source == "cold_storage":
|
||||
return CleanTemplateText(ev.Input + " " + ev.Response)
|
||||
return CleanText(ev.Input + " " + ev.Response)
|
||||
default:
|
||||
return CleanTemplateText(ev.Input)
|
||||
return CleanText(ev.Input)
|
||||
}
|
||||
}
|
||||
|
||||
@ -275,11 +275,11 @@ func TestRealContextEmbedderStats(t *testing.T) {
|
||||
var text string
|
||||
switch {
|
||||
case ev.Source == "agent" && ev.Response != "":
|
||||
text = CleanTemplateText(ev.Response)
|
||||
text = CleanText(ev.Response)
|
||||
case ev.Source == "cold_storage":
|
||||
text = CleanTemplateText(ev.Input + " " + ev.Response)
|
||||
text = CleanText(ev.Input + " " + ev.Response)
|
||||
default:
|
||||
text = CleanTemplateText(ev.Input)
|
||||
text = CleanText(ev.Input)
|
||||
}
|
||||
vec := e.Vectorize(text)
|
||||
origLen := len(ev.Input + ev.Response)
|
||||
|
||||
@ -54,6 +54,11 @@ func RegisterFactory(name string, factory NativeFactory) {
|
||||
globalFactories.Store(name, factory)
|
||||
}
|
||||
|
||||
// PluginToolCleaner 定义插件工具注销接口,由 StageHost 实现。
|
||||
type PluginToolCleaner interface {
|
||||
UnregisterPluginTools(pluginName string)
|
||||
}
|
||||
|
||||
type Registry struct {
|
||||
mu sync.RWMutex
|
||||
plugins map[string]sdk.Plugin
|
||||
@ -75,6 +80,11 @@ type Registry struct {
|
||||
regTool sdk.ToolRegistrar
|
||||
regStage sdk.StageRegistrar
|
||||
regAPI sdk.APIRegistrar
|
||||
|
||||
toolCleaner PluginToolCleaner
|
||||
|
||||
knownDisabled map[string]bool
|
||||
textCleaners []func(string) string
|
||||
}
|
||||
|
||||
func NewRegistry() *Registry {
|
||||
@ -82,6 +92,7 @@ func NewRegistry() *Registry {
|
||||
plugins: make(map[string]sdk.Plugin),
|
||||
factories: make(map[string]NativeFactory),
|
||||
pluginAutoRestart: make(map[string]bool),
|
||||
knownDisabled: make(map[string]bool),
|
||||
}
|
||||
}
|
||||
|
||||
@ -97,6 +108,18 @@ func (r *Registry) SetPluginDir(dir string) { r.plgDir = di
|
||||
func (r *Registry) SetToolRegistrar(fn sdk.ToolRegistrar) { r.regTool = fn }
|
||||
func (r *Registry) SetStageRegistrar(fn sdk.StageRegistrar) { r.regStage = fn }
|
||||
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()
|
||||
@ -217,6 +240,13 @@ func (r *Registry) Load(dir string) error {
|
||||
if loaded[name] {
|
||||
continue
|
||||
}
|
||||
if r.isDisabled(name) {
|
||||
log.Printf("[plugin] %s is disabled, skipping", name)
|
||||
r.mu.Lock()
|
||||
r.knownDisabled[name] = true
|
||||
r.mu.Unlock()
|
||||
continue
|
||||
}
|
||||
plgDir := filepath.Join(dir, name)
|
||||
os.MkdirAll(plgDir, 0755)
|
||||
|
||||
@ -240,6 +270,7 @@ 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)
|
||||
}
|
||||
@ -247,7 +278,26 @@ func (r *Registry) Load(dir string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Registry) isDisabled(name string) bool {
|
||||
if r.cfgReg == nil {
|
||||
return false
|
||||
}
|
||||
v, err := r.cfgReg.PluginConfig(name).Get("disabled")
|
||||
if err != nil || v == nil {
|
||||
return false
|
||||
}
|
||||
return fmt.Sprint(v) == "true"
|
||||
}
|
||||
|
||||
func (r *Registry) loadOne(plgDir, name string) bool {
|
||||
if r.isDisabled(name) {
|
||||
log.Printf("[plugin] %s is disabled, skipping", name)
|
||||
r.mu.Lock()
|
||||
r.knownDisabled[name] = true
|
||||
r.mu.Unlock()
|
||||
return false
|
||||
}
|
||||
|
||||
// 1) 查找工厂(init 自注册或 RegisterNative)
|
||||
r.mu.RLock()
|
||||
factory, hasFactory := r.factories[name]
|
||||
@ -303,6 +353,7 @@ 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
|
||||
@ -319,6 +370,7 @@ 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) {
|
||||
@ -382,6 +434,94 @@ func (r *Registry) AutoRestartEnabled(name string) bool {
|
||||
return enabled
|
||||
}
|
||||
|
||||
func (r *Registry) IsDisabled(name string) bool {
|
||||
return r.isDisabled(name)
|
||||
}
|
||||
|
||||
func (r *Registry) Enable(name string) error {
|
||||
if r.cfgReg != nil {
|
||||
r.cfgReg.PluginConfig(name).Set("disabled", "false")
|
||||
}
|
||||
r.mu.Lock()
|
||||
delete(r.knownDisabled, name)
|
||||
r.mu.Unlock()
|
||||
plgDir := filepath.Join(r.plgDir, name)
|
||||
if r.loadOne(plgDir, name) {
|
||||
log.Printf("[plugin] enabled: %s", name)
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("enable plugin %s failed", name)
|
||||
}
|
||||
|
||||
func (r *Registry) Disable(name string) error {
|
||||
r.mu.Lock()
|
||||
p, ok := r.plugins[name]
|
||||
if ok {
|
||||
if err := p.Stop(); err != nil {
|
||||
log.Printf("[plugin] stop %s for disable: %v", name, err)
|
||||
}
|
||||
delete(r.plugins, name)
|
||||
for i, inst := range r.instances {
|
||||
if inst.Name() == name {
|
||||
r.instances = append(r.instances[:i], r.instances[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
r.knownDisabled[name] = true
|
||||
r.mu.Unlock()
|
||||
|
||||
if r.toolCleaner != nil {
|
||||
r.toolCleaner.UnregisterPluginTools(name)
|
||||
}
|
||||
|
||||
if r.cfgReg != nil {
|
||||
r.cfgReg.PluginConfig(name).Set("disabled", "true")
|
||||
}
|
||||
log.Printf("[plugin] disabled: %s", name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListKnown 返回所有已知插件(已加载 + 已禁用 + 已安装但未加载)。
|
||||
func (r *Registry) ListKnown() []string {
|
||||
r.mu.RLock()
|
||||
known := make(map[string]bool)
|
||||
for name := range r.plugins {
|
||||
known[name] = true
|
||||
}
|
||||
for name := range r.knownDisabled {
|
||||
known[name] = true
|
||||
}
|
||||
r.mu.RUnlock()
|
||||
|
||||
if r.plgDir != "" {
|
||||
entries, _ := os.ReadDir(r.plgDir)
|
||||
for _, e := range entries {
|
||||
if e.IsDir() {
|
||||
known[e.Name()] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
r.mu.RLock()
|
||||
for name := range r.factories {
|
||||
known[name] = true
|
||||
}
|
||||
r.mu.RUnlock()
|
||||
|
||||
globalFactories.Range(func(key, val interface{}) bool {
|
||||
known[key.(string)] = true
|
||||
return true
|
||||
})
|
||||
|
||||
list := make([]string, 0, len(known))
|
||||
for name := range known {
|
||||
list = append(list, name)
|
||||
}
|
||||
sort.Strings(list)
|
||||
return list
|
||||
}
|
||||
|
||||
func (r *Registry) PluginMetas() map[string]PluginMeta {
|
||||
metas := make(map[string]PluginMeta)
|
||||
globalPluginMeta.Range(func(key, val interface{}) bool {
|
||||
|
||||
@ -197,6 +197,7 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
|
||||
s.RegisterTool("terminal_create", sdk.ToolDef{
|
||||
Name: "terminal_create",
|
||||
Description: "创建一个新的交互式终端会话。返回终端 ID,后续通过此 ID 进行读写操作。适用于运行交互式程序如 vim、ssh、top、nano 等。终端默认 5 分钟后自动关闭,可通过 timeout 参数调整。",
|
||||
NoMemory: true,
|
||||
Parameters: map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{
|
||||
@ -225,6 +226,7 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
|
||||
s.RegisterTool("terminal_write", sdk.ToolDef{
|
||||
Name: "terminal_write",
|
||||
Description: "向指定终端发送输入。支持普通文本和特殊键(通过 key 参数传入)。特殊键包括:enter, tab, escape, ctrl_a~ctrl_z, alt_a~alt_z, f1~f12, up, down, left, right, home, end, backspace, delete, page_up, page_down。普通文本传入 input 参数即可。",
|
||||
NoMemory: true,
|
||||
Parameters: map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{
|
||||
@ -250,6 +252,7 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
|
||||
s.RegisterTool("terminal_read", sdk.ToolDef{
|
||||
Name: "terminal_read",
|
||||
Description: "读取指定终端的当前屏幕内容。返回自上次读取以来的新输出。如需持续监控请多次调用。",
|
||||
NoMemory: true,
|
||||
Parameters: map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{
|
||||
@ -271,6 +274,7 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
|
||||
s.RegisterTool("terminal_resize", sdk.ToolDef{
|
||||
Name: "terminal_resize",
|
||||
Description: "调整指定终端的尺寸(行数和列数)。",
|
||||
NoMemory: true,
|
||||
Parameters: map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{
|
||||
@ -296,6 +300,7 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
|
||||
s.RegisterTool("terminal_close", sdk.ToolDef{
|
||||
Name: "terminal_close",
|
||||
Description: "关闭指定终端会话。释放资源。",
|
||||
NoMemory: true,
|
||||
Parameters: map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{
|
||||
@ -313,6 +318,7 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
|
||||
s.RegisterTool("terminal_list", sdk.ToolDef{
|
||||
Name: "terminal_list",
|
||||
Description: "列出所有活跃的终端会话及其状态。",
|
||||
NoMemory: true,
|
||||
Parameters: map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{},
|
||||
|
||||
@ -110,6 +110,7 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
|
||||
s.RegisterTool("cmd_run", sdk.ToolDef{
|
||||
Name: "cmd_run",
|
||||
Description: "执行一条系统命令并返回输出。适用于查询系统信息、运行脚本、操作文件等单次命令场景。命令在临时 shell 中执行,不支持交互。如需交互式终端(如 vim、ssh、top),请使用 terminal_create 相关工具。",
|
||||
NoMemory: true,
|
||||
Parameters: map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{
|
||||
|
||||
Reference in New Issue
Block a user