fix: correct context pruning order and vector alignment

- Prune context BEFORE processing (LSTM forget gate pattern)
  so LLM only sees relevant context, instead of pruning after the fact
- Fix ensureTrained() to recompute all event vectors after retraining
  vectorizer, fixing feature-space mismatch between stored vectors and
  query vector that made relevance scoring effectively random
- Add read lock to knowledge BuildTree() (data race fix)
- Log writeIndex() errors instead of discarding them
- Fix TOCTOU race in document ContextToDoc() dedup
- Fix healthcheck timing (measure elapsed before cleanup)
- Refactor waiter CLI into separate files (state, conn, config, editor,
  history, builtin) for maintainability
- Add CLI plugin API key authentication
This commit is contained in:
root
2026-07-07 17:07:38 +08:00
parent 1806a49b8b
commit bf982030aa
13 changed files with 1139 additions and 602 deletions

View File

@ -2,6 +2,7 @@ package core
import (
"context"
"errors"
"fmt"
"log"
"strings"
@ -251,16 +252,27 @@ func (a *Agent) interceptLoop() {
a.llmMu.Unlock()
if hasActiveLLM {
select {
case a.interceptCh <- clone:
default:
log.Printf("[agent] intercept channel full, queuing input for %s", evt.Source)
// 后台整理任务被打断:中断消息重新注入为独立输入(consolidation 的 process 不会路由回复)
if a.currentOutputChannel == "_consolidation_" {
log.Printf("[agent] consolidation interrupted, re-injecting input for %s/%s", evt.Source, evt.OutputChannel)
a.io.InjectInputTo(evt.Source, evt.OutputChannel, "text", map[string]interface{}{
"content": text,
"interrupt": true,
"interrupt_source": evt.Source,
"interrupt_channel": evt.OutputChannel,
})
} else {
select {
case a.interceptCh <- clone:
default:
log.Printf("[agent] intercept channel full, queuing input for %s", evt.Source)
a.io.InjectInputTo(evt.Source, evt.OutputChannel, "text", map[string]interface{}{
"content": text,
"interrupt": true,
"interrupt_source": evt.Source,
"interrupt_channel": evt.OutputChannel,
})
}
}
} else {
a.io.InjectInputTo(evt.Source, evt.OutputChannel, "text", map[string]interface{}{
@ -325,6 +337,12 @@ func (a *Agent) processMediaInput(evt *agentIO.InputEvent) {
blocks, fallback := a.mediaToBlocks(evt.Payload, evt.Type, evt.Source)
// 先遗忘再输入
archived := a.context.Prune(fallback, a.maxContextSize-1, a.docStore)
if archived > 0 {
log.Printf("[agent] pruned %d low-relevance events to document memory", archived)
}
a.context.Append(ContextEvent{
Timestamp: start,
Source: evt.Source,
@ -371,11 +389,6 @@ func (a *Agent) processMediaInput(evt *agentIO.InputEvent) {
ToolsUsed: toolsUsed,
})
archived := a.context.Prune(response, a.maxContextSize, a.docStore)
if archived > 0 {
log.Printf("[agent] pruned %d low-relevance events to document memory", archived)
}
a.emitResponse(evt, response)
}
@ -466,6 +479,12 @@ func (a *Agent) processTextInput(evt *agentIO.InputEvent, input string) {
input = stageCtx.RawMessage
// 先"遗忘"再输入:用当前输入决定淘汰哪些不相关旧事件(LSTM forget gate 模式)
archived := a.context.Prune(input, a.maxContextSize-1, a.docStore)
if archived > 0 {
log.Printf("[agent] pruned %d low-relevance events to document memory", archived)
}
a.context.Append(ContextEvent{
Timestamp: start,
Source: evt.Source,
@ -492,12 +511,6 @@ func (a *Agent) processTextInput(evt *agentIO.InputEvent, input string) {
ToolsUsed: toolsUsed,
})
// 基于相关性裁剪上下文:保留与当前输入最相关的 maxContextSize 条
archived := a.context.Prune(response, a.maxContextSize, a.docStore)
if archived > 0 {
log.Printf("[agent] pruned %d low-relevance events to document memory", archived)
}
a.emitResponse(evt, response)
if !stageCtx.NoMemory {
@ -661,6 +674,14 @@ 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")
}
// 用户对话被打断:继续下一轮 drain 打断消息,注入到当前对话上下文
continue
}
return "", toolsUsed, fmt.Errorf("all %d providers failed, last error: %w",
len(providers), llmErr)
}
@ -1159,7 +1180,11 @@ func (a *Agent) executeKnowledgeTool(tc agentAPI.ToolCall) string {
if i >= topK {
break
}
parts = append(parts, fmt.Sprintf("[%s]\n%s", k.Name, truncateStr(k.Content, 200)))
label := k.Name
if k.Category != "" {
label = k.Category + "/" + k.Name
}
parts = append(parts, fmt.Sprintf("[%s]\n%s", label, truncateStr(k.Content, 200)))
}
return strings.Join(parts, "\n---\n")
@ -1175,11 +1200,8 @@ func (a *Agent) executeKnowledgeTool(tc agentAPI.ToolCall) string {
return fmt.Sprintf("知识「%s」已创建并向量化索引(%d 字符)", name, len(content))
case "knowledge_list":
names := a.knowledge.List()
if len(names) == 0 {
return "知识库为空"
}
return "知识分类: " + strings.Join(names, ", ")
tree := a.knowledge.BuildTree()
return formatTree(tree, 0)
default:
return fmt.Sprintf("未知的知识工具: %s", tc.Name)
@ -1802,9 +1824,14 @@ func (a *Agent) distillContext() {
if a.docStore == nil {
return
}
// 心跳时执行一次安全裁剪(兜底)
// 上下文的主要裁剪在 processTextInput 中基于相关性执行
_ = a.context.Len()
// 心跳时执行安全裁剪:上下文超过 maxContextSize*2 时强制归档
n := a.context.Len()
if n > a.maxContextSize*2 {
archived := a.context.Prune("", a.maxContextSize, a.docStore)
if archived > 0 {
log.Printf("[agent] distill: pruned %d low-relevance events to document memory (total=%d)", archived, n)
}
}
}
// syncGraphToDocs — 将图记忆的实体和关系注入文档记忆层
@ -2104,6 +2131,13 @@ func (a *Agent) emitMemoryCandidate(source, input, response string, toolsUsed []
func (a *Agent) processConsolidation(input string) {
start := time.Now()
a.currentOutputChannel = "_consolidation_"
// 遗忘不相关的旧事件
archived := a.context.Prune(input, a.maxContextSize-1, a.docStore)
if archived > 0 {
log.Printf("[agent] consolidation: pruned %d low-relevance events", archived)
}
a.context.Append(ContextEvent{
Timestamp: start,
Source: "system",
@ -2121,7 +2155,6 @@ func (a *Agent) processConsolidation(input string) {
Response: response,
ToolsUsed: toolsUsed,
})
_ = a.context.Prune(response, a.maxContextSize, a.docStore)
a.emitMemoryCandidate("system", input, response, toolsUsed)
log.Printf("[agent] consolidation done (%dms, tools=%v)", time.Since(start).Milliseconds(), toolsUsed)
}
@ -2585,6 +2618,45 @@ func truncateStr(s string, max int) string {
return s
}
func formatTree(node *knowledge.TreeIndex, depth int) string {
var sb strings.Builder
indent := strings.Repeat(" ", depth)
for _, child := range node.Children {
sb.WriteString(fmt.Sprintf("%s%s/\n", indent, child.Name))
if len(child.Items) > 0 {
for _, item := range child.Items {
preview := item.Preview
if len([]rune(preview)) > 60 {
preview = string([]rune(preview)[:60]) + "..."
}
tags := ""
if len(item.Tags) > 0 {
tags = " [" + strings.Join(item.Tags, ", ") + "]"
}
sb.WriteString(fmt.Sprintf("%s · %s%s\n %s\n", indent, item.Name, tags, preview))
}
}
sb.WriteString(formatTree(child, depth+1))
}
if depth > 0 && len(node.Items) > 0 {
for _, item := range node.Items {
preview := item.Preview
if len([]rune(preview)) > 60 {
preview = string([]rune(preview)[:60]) + "..."
}
tags := ""
if len(item.Tags) > 0 {
tags = " [" + strings.Join(item.Tags, ", ") + "]"
}
sb.WriteString(fmt.Sprintf(" · %s%s\n %s\n", item.Name, tags, preview))
}
}
if sb.Len() == 0 {
sb.WriteString("(空)")
}
return sb.String()
}
func (a *Agent) resolveToolPlugin(name string) string {
if a.stageHost != nil {
if plugin := a.stageHost.ToolPlugin(name); plugin != "" {

View File

@ -223,6 +223,10 @@ func (c *RelevanceContext) ensureTrained() {
texts[i] = evt.Input + " " + evt.Response
}
c.veczer.Train(texts)
// 重算所有事件向量,与新的向量化器特征空间对齐
for _, evt := range c.events {
evt.Vector = c.veczer.Vectorize(evt.Input + " " + evt.Response)
}
c.trained = true
}
}

View File

@ -1,6 +1,7 @@
package knowledge
import (
"encoding/json"
"fmt"
"log"
"os"
@ -13,17 +14,67 @@ import (
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
)
// Knowledge — 单条知识
type Knowledge struct {
Name string `json:"name"`
Content string `json:"content"`
Path string `json:"path"`
Category string `json:"category,omitempty"` // 父路径,如 "tech/go"
Tags []string `json:"tags"`
UpdatedAt time.Time `json:"updated_at"`
Meta map[string]string `json:"meta,omitempty"`
}
// Store — 知识库,文件系统 + 向量索引
// IndexItem — 索引条目,包含向量特征和内容摘要
type IndexItem struct {
Name string `json:"name"`
Preview string `json:"preview"` // 前 200 字摘要
Tags []string `json:"tags"`
Vector map[string]float64 `json:"vector"` // TF-IDF 特征向量(top-N 特征)
Size int `json:"size"` // 内容总字节数
}
// TreeIndex — 树状索引节点
type TreeIndex struct {
Name string `json:"name"`
Children map[string]*TreeIndex `json:"children,omitempty"`
Items []IndexItem `json:"items,omitempty"` // 此节点下的知识条目(含向量)
}
func newTreeIndex(name string) *TreeIndex {
return &TreeIndex{Name: name, Children: make(map[string]*TreeIndex)}
}
// compressVector 压缩向量:保留 topN 个权重最高的特征
func compressVector(v vector.Vector, topN int) map[string]float64 {
if len(v) <= topN {
out := make(map[string]float64, len(v))
for k, w := range v {
out[k] = w
}
return out
}
type kv struct {
k string
v float64
}
sorted := make([]kv, 0, len(v))
for k, w := range v {
sorted = append(sorted, kv{k, w})
}
sort.Slice(sorted, func(i, j int) bool {
return sorted[i].v > sorted[j].v
})
if topN > len(sorted) {
topN = len(sorted)
}
sorted = sorted[:topN]
out := make(map[string]float64, topN)
for _, kv := range sorted {
out[kv.k] = kv.v
}
return out
}
type Store struct {
root string
vec *vector.Store
@ -31,15 +82,17 @@ type Store struct {
mu sync.RWMutex
items map[string]*Knowledge
indexPath string
summaries []string
}
func NewStore(root string) *Store {
return &Store{
root: root,
vec: vector.NewStore(),
veczer: vector.NewTFIDFVectorizer(3),
items: make(map[string]*Knowledge),
root: root,
indexPath: filepath.Join(root, ".index.json"),
vec: vector.NewStore(),
veczer: vector.NewTFIDFVectorizer(3),
items: make(map[string]*Knowledge),
}
}
@ -50,13 +103,16 @@ func (s *Store) Start() error {
if err := s.scanAll(); err != nil {
log.Printf("[knowledge] scan error: %v", err)
}
// 重建索引文件
if err := s.writeIndex(); err != nil {
log.Printf("[knowledge] write index error: %v", err)
}
log.Printf("[knowledge] started with %d items, %d vectors", len(s.items), s.vec.Size())
return nil
}
func (s *Store) Stop() {}
// Search — 向量查询知识
func (s *Store) Search(query string, topK int) []*Knowledge {
s.mu.RLock()
defer s.mu.RUnlock()
@ -77,48 +133,56 @@ func (s *Store) Search(query string, topK int) []*Knowledge {
return out
}
// Add — 添加或更新知识
func (s *Store) Add(name, content string) error {
s.mu.Lock()
defer s.mu.Unlock()
// 创建知识目录
dir := filepath.Join(s.root, sanitize(name))
// 解析层级:将 "/" 作为路径分隔符
category := ""
leaf := name
if idx := strings.LastIndex(name, "/"); idx >= 0 {
category = name[:idx]
leaf = name[idx+1:]
}
dirName := sanitize(leaf)
if category != "" {
dirName = sanitize(category) + "/" + dirName
}
dir := filepath.Join(s.root, dirName)
if err := os.MkdirAll(dir, 0755); err != nil {
return fmt.Errorf("create knowledge dir: %w", err)
}
// 写入知识文件
path := filepath.Join(dir, "content.md")
if err := os.WriteFile(path, []byte(content), 0644); err != nil {
return fmt.Errorf("write knowledge: %w", err)
}
now := time.Now()
id := sanitize(name)
k := &Knowledge{
Name: name,
Name: id,
Content: content,
Path: path,
Category: sanitize(category),
Tags: extractKeywords(name + " " + content),
UpdatedAt: now,
}
// 生成 ID = 目录名
id := sanitize(name)
s.items[id] = k
vec := s.veczer.Vectorize(name + " " + content)
s.vec.Insert(id, name+": "+content, vec, map[string]string{
"name": name, "path": path,
})
s.summaries = append(s.summaries, name+" "+content)
if err := s.writeIndex(); err != nil {
log.Printf("[knowledge] write index error after adding %s: %v", name, err)
}
log.Printf("[knowledge] added: %s (%d bytes)", name, len(content))
return nil
}
// SearchCategories — 返回所有知识类别
func (s *Store) SearchCategories(query string, topK int) []string {
s.mu.RLock()
defer s.mu.RUnlock()
@ -157,6 +221,9 @@ func (s *Store) Remove(name string) error {
}
delete(s.items, id)
s.vec.Remove(id)
if err := s.writeIndex(); err != nil {
log.Printf("[knowledge] write index error after removing %s: %v", name, err)
}
return nil
}
@ -167,6 +234,7 @@ func (s *Store) Stats() map[string]interface{} {
"knowledge_count": len(s.items),
"vector_count": s.vec.Size(),
"root": s.root,
"index_file": s.indexPath,
}
}
@ -181,6 +249,89 @@ func (s *Store) List() []string {
return names
}
// BuildTree 从当前知识库构建树状索引(含向量特征)
func (s *Store) BuildTree() *TreeIndex {
s.mu.RLock()
defer s.mu.RUnlock()
root := newTreeIndex("root")
for _, k := range s.items {
node := root
if k.Category != "" {
parts := strings.Split(k.Category, "/")
for _, part := range parts {
if part == "" {
continue
}
if _, ok := node.Children[part]; !ok {
node.Children[part] = newTreeIndex(part)
}
node = node.Children[part]
}
}
// 获取该条目的向量并压缩
vec := s.veczer.Vectorize(k.Name + " " + k.Content)
preview := []rune(k.Content)
previewStr := ""
if len(preview) > 200 {
previewStr = string(preview[:200]) + "..."
} else {
previewStr = string(preview)
}
item := IndexItem{
Name: k.Name,
Preview: previewStr,
Tags: k.Tags,
Vector: compressVector(vec, 20),
Size: len(k.Content),
}
node.Items = append(node.Items, item)
}
return root
}
// SearchTree 树状搜索:在树节点下搜索,返回按分类聚合的结果
func (s *Store) SearchTree(query string, topK int) map[string][]*Knowledge {
s.mu.RLock()
defer s.mu.RUnlock()
if topK <= 0 {
topK = 10
}
vec := s.veczer.Vectorize(query)
results := s.vec.Search(vec, topK*2)
categorized := make(map[string][]*Knowledge)
for _, r := range results {
if k, ok := s.items[r.ID]; ok {
cat := k.Category
if cat == "" {
cat = "未分类"
}
categorized[cat] = append(categorized[cat], k)
}
}
out := make(map[string][]*Knowledge)
for cat, items := range categorized {
if len(items) > topK {
items = items[:topK]
}
out[cat] = items
}
return out
}
// writeIndex 写入 .index.json 树状索引文件(含向量和摘要)
func (s *Store) writeIndex() error {
tree := s.BuildTree()
data, err := json.MarshalIndent(tree, "", " ")
if err != nil {
return err
}
return os.WriteFile(s.indexPath, data, 0644)
}
// ——— internal ———
func (s *Store) scanAll() error {
@ -193,34 +344,17 @@ func (s *Store) scanAll() error {
if !entry.IsDir() {
continue
}
dir := filepath.Join(s.root, entry.Name())
contentPath := filepath.Join(dir, "content.md")
data, err := os.ReadFile(contentPath)
if err != nil {
// skip hidden dirs
if strings.HasPrefix(entry.Name(), ".") {
continue
}
name := entry.Name()
content := string(data)
now := time.Now()
k := &Knowledge{
Name: name,
Content: content,
Path: contentPath,
Tags: extractKeywords(name + " " + content),
UpdatedAt: now,
}
s.items[name] = k
s.summaries = append(s.summaries, name+" "+content)
s.scanDir("", entry.Name())
}
// 训练向量化器
if len(s.summaries) > 0 {
s.veczer.Train(s.summaries)
}
// 构建向量索引
for _, k := range s.items {
vec := s.veczer.Vectorize(k.Name + " " + k.Content)
s.vec.Insert(k.Name, k.Name+": "+k.Content, vec, map[string]string{
@ -231,11 +365,48 @@ func (s *Store) scanAll() error {
return nil
}
// scanDir 递归扫描目录
// category: 父级路径(从知识库根目录算起),如 "tech/go"
// dirName: 当前目录相对路径(从知识库根目录算起)
func (s *Store) scanDir(category, dirName string) {
dir := filepath.Join(s.root, dirName)
contentPath := filepath.Join(dir, "content.md")
data, err := os.ReadFile(contentPath)
if err == nil {
name := dirName
content := string(data)
now := time.Now()
k := &Knowledge{
Name: name,
Content: content,
Path: contentPath,
Category: category,
Tags: extractKeywords(dirName + " " + content),
UpdatedAt: now,
}
s.items[name] = k
s.summaries = append(s.summaries, name+" "+content)
return
}
// 无 content.md => 是分类目录,递归子目录
subEntries, err := os.ReadDir(dir)
if err != nil {
return
}
for _, sub := range subEntries {
if !sub.IsDir() || strings.HasPrefix(sub.Name(), ".") {
continue
}
childDir := dirName + "/" + sub.Name()
s.scanDir(dirName, childDir)
}
}
func sanitize(name string) string {
name = strings.ToLower(name)
name = strings.TrimSpace(name)
name = strings.ReplaceAll(name, " ", "_")
name = strings.ReplaceAll(name, "/", "_")
name = strings.ReplaceAll(name, "\\", "_")
return name
}
@ -255,10 +426,8 @@ func extractKeywords(text string) []string {
var keywords []string
runes := []rune(text)
seen := make(map[string]bool)
// bi-gram
for i := 0; i < len(runes)-1; i++ {
word := string(runes[i : i+2])
if !stopWords[word] && len(strings.TrimSpace(word)) == len(word) && !seen[word] {

View File

@ -93,7 +93,7 @@ func (s *Store) Insert(doc *Doc) error {
return nil
}
// ContextToDoc — 将一段上下文对话历史提炼为文档
// ContextToDoc — 将一段上下文对话历史提炼为文档(带内容去重)
func (s *Store) ContextToDoc(source string, entries []ContextEntry) (*Doc, error) {
if len(entries) == 0 {
return nil, nil
@ -108,26 +108,46 @@ func (s *Store) ContextToDoc(source string, entries []ContextEntry) (*Doc, error
parts = append(parts, line)
}
content := strings.Join(parts, "\n")
contentHash := simpleHash(content)
// 去重:检查是否已有相同 hash 的文档(在锁内完成创建/更新)
summary := summarizeEntries(entries)
tags := extractTags(entries)
entities := extractEntities(entries)
s.mu.Lock()
for _, d := range s.docs {
if d.Meta != nil && d.Meta["content_hash"] == contentHash {
d.UpdatedAt = time.Now()
d.LastAccess = time.Now()
d.Content = content
d.Source = source
d.Summary = summary
d.Tags = tags
d.Entities = entities
s.dirty = true
s.mu.Unlock()
return d, nil
}
}
id := fmt.Sprintf("doc_%d", time.Now().UnixNano())
doc := &Doc{
ID: fmt.Sprintf("doc_%d", time.Now().UnixNano()),
Summary: summary,
Content: content,
Tags: tags,
Entities: entities,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
Source: source,
ID: id,
Summary: summary,
Content: content,
Tags: tags,
Entities: entities,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
LastAccess: time.Now(),
AccessCount: 1,
Source: source,
Meta: map[string]string{"content_hash": contentHash},
}
if err := s.Insert(doc); err != nil {
return nil, err
}
s.docs[id] = doc
s.dirty = true
s.mu.Unlock()
return doc, nil
}
@ -422,3 +442,12 @@ func truncate(s string, max int) string {
}
return s
}
func simpleHash(s string) string {
// 简单的基于内容的哈希,用于去重
h := 0
for _, r := range s {
h = h*31 + int(r)
}
return fmt.Sprintf("h%08x", h)
}

View File

@ -108,6 +108,26 @@ func (p *Plugin) handleConn(conn net.Conn, s *sdk.PluginSDK) {
defer p.wg.Done()
scanner := bufio.NewScanner(conn)
apiKey := p.webuiAPIKey()
if apiKey != "" {
if !scanner.Scan() {
return
}
line := scanner.Text()
if !strings.HasPrefix(line, "/auth ") || strings.TrimSpace(line[6:]) != apiKey {
writeLine(conn, map[string]interface{}{
"type": "error",
"error": "unauthorized",
})
return
}
writeLine(conn, map[string]interface{}{
"type": "response",
"content": "authenticated",
})
}
for scanner.Scan() {
line := scanner.Text()
if line == "" {
@ -136,6 +156,19 @@ func (p *Plugin) handleConn(conn net.Conn, s *sdk.PluginSDK) {
}
}
func (p *Plugin) webuiAPIKey() string {
if cfgReg == nil {
return ""
}
ps := cfgReg.PluginConfig("webui")
if v, _ := ps.Get("api_key"); v != nil {
if s, ok := v.(string); ok {
return s
}
}
return ""
}
func (p *Plugin) handleBuiltin(conn net.Conn, line string, s *sdk.PluginSDK) bool {
parts := strings.Fields(line)
if len(parts) == 0 {

View File

@ -456,8 +456,12 @@ func (p *Plugin) testKnowledgeRaw() checkResult {
}
results := hcKnowledge.Search("健康检查测试标记", 3)
elapsed := time.Since(start)
// 清理测试条目,避免积累
hcKnowledge.Remove(marker)
if len(results) > 0 {
return checkResult{
Name: "knowledge",