feat: 完整实现 NLP 三元组提取系统 + token budget 上下文分配

- 重写 extractor.go: 分句、17条 POS 模板、依存模板 + COO 链、ATT合并
- parser.go: 分句循环 + TransE 向量验证(h+r≈t)
- fallback.go: jieba POS 降级解析器
- bridge.go: nlp.Triple ↔ memory.Triple 转换
- pipeline.go: extractKeyTriples 改用 NLP 提取器, 删除5条旧前缀规则
- distill.go: docToTriples 改用 NLP 提取器
- reorgGraph: 语义相似度增强检测, 保持纯 LLM 决断
- Provider 接口加 MaxContextTokens() + 模型窗口映射表
- tokenbudget.go: 中文 token 估算器 + budget 分配(80%利用率)
- process.go/buildSystemPrompt: 按 token 预算截断 memory+timeline
This commit is contained in:
root
2026-07-27 15:26:23 +08:00
parent d19b7bd13e
commit 1cb3e87dde
30 changed files with 1508 additions and 773 deletions

View File

@ -105,6 +105,8 @@ type Agent struct {
noMergeMarkers map[string]int
noMergeMu sync.Mutex
// 词嵌入模型,用于实体语义相似度计算
embedder *memory.StaticEmbedder
}
type AgentConfig struct {
@ -187,6 +189,7 @@ func New(cfg AgentConfig) *Agent {
pluginHealth: newPluginHealthTracker(),
thinkingEnabled: cfg.ThinkingEnabled,
inputCfg: cfg.InputProcessing,
embedder: embedder,
noMergeMarkers: make(map[string]int),
}

View File

@ -35,14 +35,11 @@ func TestDocToTriples(t *testing.T) {
triples := docToTriples(doc)
foundSummary := false
foundRel := false
foundSource := false
for _, tr := range triples {
switch {
case tr.Subject == "文档" && tr.Relation == "主题":
foundSummary = true
case tr.Relation == "关联":
foundRel = true
case tr.Subject == "文档" && tr.Relation == "来源":
foundSource = true
}
@ -54,9 +51,6 @@ func TestDocToTriples(t *testing.T) {
if !foundSource {
t.Error("missing '来源' triple")
}
if needJieba() && !foundRel {
t.Error("missing '关联' triple with jieba available")
}
}
func TestDocToTriplesNil(t *testing.T) {
@ -88,7 +82,6 @@ func TestDocToTriplesTypes(t *testing.T) {
triples := docToTriples(doc)
// 主题 and 来源 triples have Subject=文档
for _, tr := range triples {
if tr.Subject == "文档" {
if tr.SubjectType != "Concept" {
@ -97,14 +90,6 @@ func TestDocToTriplesTypes(t *testing.T) {
if tr.Confidence != 1.0 {
t.Errorf("文档 triple confidence should be 1.0, got %f", tr.Confidence)
}
} else {
// 关联 triples use extracted terms as subject/object
if tr.Relation != "关联" {
t.Errorf("non-文档 triple should have 关联 relation, got %q", tr.Relation)
}
if tr.Confidence != 0.8 {
t.Errorf("关联 triple confidence should be 0.8, got %f", tr.Confidence)
}
}
// all should have SubjectType/ObjectType set
if tr.SubjectType == "" || tr.ObjectType == "" {

View File

@ -40,11 +40,8 @@ func TestDocToTriplesConversation(t *testing.T) {
}
triples := docToTriples(doc)
minLen := 2
hasJieba := needJieba()
if hasJieba && len(triples) <= minLen {
t.Errorf("expected more than %d triples with jieba, got %d", minLen, len(triples))
if len(triples) < 2 {
t.Errorf("expected at least 2 triples (主题+来源), got %d", len(triples))
}
for i, tr := range triples {
@ -55,19 +52,6 @@ func TestDocToTriplesConversation(t *testing.T) {
t.Errorf("triple[%d] has non-positive confidence: %+v", i, tr)
}
}
relCount := 0
for _, tr := range triples {
if tr.Relation == "关联" {
relCount++
if tr.Subject == tr.Object {
t.Errorf("关联 triple has same subject and object: %+v", tr)
}
}
}
if hasJieba && relCount == 0 {
t.Errorf("expected 关联 triples with jieba enabled, got 0 in %+v", triples)
}
}
func TestDocToTriplesMultiLine(t *testing.T) {

View File

@ -4,12 +4,13 @@ import (
"fmt"
"log"
"runtime/debug"
"strings"
"time"
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/document"
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
"gitcode.com/JianFeeeee/HomeAgent/internal/nlp"
)
type ConsolidationTask struct {
@ -107,15 +108,18 @@ func (a *Agent) reorgGraph() {
return
}
llmCandidates := 0
maxCandidates := 5
candidates := 0
for i := 0; i < len(result.Entities) && candidates < maxCandidates; i++ {
for j := i + 1; j < len(result.Entities) && candidates < maxCandidates; j++ {
for i := 0; i < len(result.Entities) && llmCandidates < maxCandidates; i++ {
for j := i + 1; j < len(result.Entities) && llmCandidates < maxCandidates; j++ {
ea, eb := result.Entities[i].Name, result.Entities[j].Name
if ea > eb {
ea, eb = eb, ea
}
key := ea + "||" + eb
// 跳过已标记"不合并"的实体对
a.noMergeMu.Lock()
rounds, ok := a.noMergeMarkers[key]
if ok {
@ -130,9 +134,16 @@ func (a *Agent) reorgGraph() {
if ok {
continue
}
// 复合相似度:字符二元组 + 语义向量(仅增强检测,不做自动合并)
sim := entitySimilarity(result.Entities[i].Name, result.Entities[j].Name)
semSim := entitySemanticSimilarity(result.Entities[i].Name, result.Entities[j].Name, a.embedder)
if semSim > sim {
sim = semSim
}
if sim > 0.75 {
candidates++
llmCandidates++
a.enqueueConsolidationTask(ConsolidationTask{
Type: "entity_merge",
Reason: fmt.Sprintf(
@ -155,83 +166,24 @@ func (a *Agent) reorgGraph() {
}
}
if candidates > 0 {
log.Printf("[agent] graph reorg: %d merge candidates sent for LLM decision", candidates)
if llmCandidates > 0 {
log.Printf("[agent] graph reorg: %d merge candidates sent for LLM decision", llmCandidates)
} else {
log.Printf("[agent] graph reorg: no similar entities found")
}
a.evaluateGraphQuality()
}
func (a *Agent) evaluateGraphQuality() {
if a.memory == nil {
return
// entitySemanticSimilarity 使用词嵌入向量余弦相似度计算实体名语义相似度
func entitySemanticSimilarity(a, b string, embedder *memory.StaticEmbedder) float64 {
if a == "" || b == "" || embedder == nil || !embedder.Loaded() {
return 0
}
pending, err := a.memory.RecallPending(10)
if err != nil {
log.Printf("[agent] recall pending relations error: %v", err)
return
va := embedder.Vectorize(a)
vb := embedder.Vectorize(b)
if len(va) == 0 || len(vb) == 0 {
return 0
}
if len(pending) == 0 {
return
}
var lowQuality []string
var pendingIDs []int64
var skipIDs []int64
for _, r := range pending {
isLow := false
if (r.SourceName == "用户" || r.SourceName == "AI") &&
(r.RelationType == "提及" || r.RelationType == "回应") {
isLow = true
} else if r.RelationType == "关联" {
isLow = true
} else if r.Confidence < 0.3 && r.RelationType != "" {
isLow = true
}
if !isLow {
skipIDs = append(skipIDs, r.ID)
continue
}
pendingIDs = append(pendingIDs, r.ID)
label := fmt.Sprintf("「%s」-「%s」→「%s」", r.SourceName, r.RelationType, r.TargetName)
if r.RelationType == "关联" {
label += "(jieba 共现)"
} else if r.Confidence < 0.3 {
label += fmt.Sprintf("(confidence=%.1f)", r.Confidence)
}
lowQuality = append(lowQuality, label)
}
if len(skipIDs) > 0 {
a.memory.UpdateEvalStatusBatch(skipIDs, "approved")
}
if len(lowQuality) == 0 {
return
}
if err := a.memory.UpdateEvalStatusBatch(pendingIDs, "evaluating"); err != nil {
log.Printf("[agent] mark relations evaluating error: %v", err)
return
}
a.enqueueConsolidationTask(ConsolidationTask{
Type: "graph_quality",
Reason: fmt.Sprintf(
"图数据库中发现 %d 条低质量关系,请逐条判断是否应该删除(保留 = keep,删除 = discard):\n%s",
len(lowQuality),
strings.Join(lowQuality, "\n"),
),
Data: map[string]interface{}{
"candidates": lowQuality,
"action": "evaluate_quality",
},
})
log.Printf("[agent] graph quality: %d pending relations sent for LLM evaluation", len(lowQuality))
return vector.CosineSimilarity(va, vb)
}
func entitySimilarity(a, b string) float64 {
@ -286,6 +238,7 @@ func docToTriples(doc *document.Doc) []memory.Triple {
return nil
}
// 文档元数据
triples = append(triples, memory.Triple{
Subject: "文档",
SubjectType: "Concept",
@ -295,22 +248,15 @@ func docToTriples(doc *document.Doc) []memory.Triple {
Confidence: 1.0,
})
lines := strings.Split(doc.Content, "\n")
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" {
continue
}
terms := memory.CutExact(line)
for i := 0; i < len(terms)-1; i++ {
triples = append(triples, memory.Triple{
Subject: terms[i],
SubjectType: "Concept",
Relation: "关联",
Object: terms[i+1],
ObjectType: "Concept",
Confidence: 0.8,
})
// NLP 通用提取
e := nlp.NewExtractor(nil)
result := e.Extract(doc.Content)
if result != nil {
for _, nt := range result.Triples {
mt := nlp.ToMemoryTriple(nt)
if mt.Subject != "" && mt.Relation != "" && mt.Object != "" {
triples = append(triples, mt)
}
}
}
@ -354,13 +300,5 @@ func (a *Agent) processConsolidation(evt *agentIO.InputEvent, input string) {
return
}
if a.memory != nil {
if n, err := a.memory.ResolveEvaluating(); err != nil {
log.Printf("[agent] resolve evaluating relations error: %v", err)
} else if n > 0 {
log.Printf("[agent] resolved %d evaluating relations to approved", n)
}
}
log.Printf("[agent] consolidation done (%dms, tools=%v)", time.Since(start).Milliseconds(), toolsUsed)
}

View File

@ -20,18 +20,21 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri
return "", nil, nil, fmt.Errorf("agent: no LLM provider configured")
}
memContext := a.buildMemoryContext(input)
budget := ComputeTokenBudget(a.provider, a.systemPrompt)
memContext := a.buildMemoryContext(input, budget.MemoryTokens)
sysPrompt := a.buildSystemPrompt(memContext, input)
tools := a.buildToolDefs()
msgs := a.buildMessages(sysPrompt, input)
msgs := a.buildMessages(sysPrompt, input, budget.ContextTokens)
if blocks, ok := stageCtx.Extra["media_blocks"].([]agentAPI.ContentBlock); ok && len(blocks) > 0 {
if len(msgs) > 0 {
msgs[len(msgs)-1].Blocks = blocks
}
}
log.Printf("[agent] tool call loop start, %d tools, %d context events, personality=%t, docs=%d",
log.Printf("[agent] tool call loop start, max_ctx=%d target=%d fixed=%d mem=%d ctx=%d %d tools, %d events, personality=%t, docs=%d",
budget.MaxContext, budget.TargetUsage, budget.FixedTokens, budget.MemoryTokens, budget.ContextTokens,
len(tools), a.context.Len(),
a.personality != nil && a.personality.Content != "",
a.docStoreSize())
@ -295,7 +298,7 @@ func (a *Agent) docStoreSize() int {
return 0
}
func (a *Agent) formatMergedTimeline() string {
func (a *Agent) formatMergedTimeline(maxTokens int) string {
a.context.mu.Lock()
events := make([]*ContextEvent, len(a.context.events))
copy(events, a.context.events)
@ -305,9 +308,35 @@ func (a *Agent) formatMergedTimeline() string {
return ""
}
// 第一轮:从最新到最旧,计算在预算内能放多少条
headerTokens := EstimateTokens("【对话时序】\n")
remaining := maxTokens - headerTokens
include := 0
for i := len(events) - 1; i >= 0; i-- {
e := events[i]
est := len(e.Source) + len(e.Input) + 40
if e.Response != "" {
est += 120
}
estTokens := est * 2
if remaining-estTokens < 0 && include > 0 {
break
}
remaining -= estTokens
include++
}
if include == 0 && len(events) > 0 {
include = 1
}
// 第二轮:按时间正序渲染
start := len(events) - include
if start < 0 {
start = 0
}
var sb strings.Builder
sb.WriteString("【对话时序】\n")
for _, e := range events {
for _, e := range events[start:] {
sb.WriteString(fmt.Sprintf("[%s] %s: %s",
e.Timestamp.Format("15:04:05"), e.Source, e.Input))
if len(e.ToolsUsed) > 0 {
@ -321,11 +350,11 @@ func (a *Agent) formatMergedTimeline() string {
return sb.String()
}
func (a *Agent) buildMessages(sysPrompt, input string) []agentAPI.Message {
func (a *Agent) buildMessages(sysPrompt, input string, ctxTokens int) []agentAPI.Message {
msgs := []agentAPI.Message{{Role: "system", Content: sysPrompt}}
if ctxStr := a.formatMergedTimeline(); ctxStr != "" {
msgs = append(msgs, agentAPI.Message{Role: "system", Content: ctxStr})
if ctxTok := a.formatMergedTimeline(ctxTokens); ctxTok != "" {
msgs = append(msgs, agentAPI.Message{Role: "system", Content: ctxTok})
}
msgs = append(msgs, agentAPI.Message{Role: "user", Content: input})

View File

@ -0,0 +1,86 @@
package core
import (
"unicode/utf8"
"gitcode.com/JianFeeeee/HomeAgent/internal/agent/api"
)
// TokenBudget 上下文 token 预算分配结果
type TokenBudget struct {
MaxContext int // 模型窗口上限
TargetUsage int // 目标使用量(max * utilizationRate)
FixedTokens int // 固定部分(system prompt base + tools + rules)
MemoryTokens int // memory context 可用预算
ContextTokens int // 上下文事件可用预算
Reserved int // 预留(response 空间)
}
// EstimateTokens 粗略估算 token 数
// 中文 ~1.5 token/字,英文 ~0.3 token/字符
// 保守估计取 max(1, runeCount * 2),对混合文本足够安全
func EstimateTokens(text string) int {
if text == "" {
return 0
}
runeCount := utf8.RuneCountInString(text)
if runeCount == 0 {
return 0
}
t := runeCount * 2
if t < 1 {
return 1
}
return t
}
// ComputeTokenBudget 计算各部分的 token 预算
// utilizationRate 为目标窗口利用率(0.0-1.0),预留 1-utilizationRate 给 response
// 固定部分优先保障,剩余预算 1:2 分配给 memory context 和 context events
func ComputeTokenBudget(provider api.Provider, systemPromptBase string) TokenBudget {
maxCtx := provider.MaxContextTokens()
if maxCtx <= 0 {
maxCtx = 32768
}
utilizationRate := 0.8
targetUsage := int(float64(maxCtx) * utilizationRate)
reserved := maxCtx - targetUsage
fixedTokens := EstimateTokens(systemPromptBase)
available := targetUsage - fixedTokens
if available < 0 {
available = 0
}
// memory context 占 1/3,context events 占 2/3
memTokens := available / 3
ctxTokens := available - memTokens
return TokenBudget{
MaxContext: maxCtx,
TargetUsage: targetUsage,
FixedTokens: fixedTokens,
MemoryTokens: memTokens,
ContextTokens: ctxTokens,
Reserved: reserved,
}
}
// TruncateByTokens 截断字符串至不超过 maxTokens 估计值
func TruncateByTokens(s string, maxTokens int) string {
if maxTokens <= 0 || s == "" {
return ""
}
runes := []rune(s)
if len(runes)*2 <= maxTokens {
return s
}
// 从开头保留 maxTokens/2 个字符(每个字符约 2 token)
keep := maxTokens / 2
if keep >= len(runes) {
return s
}
return string(runes[:keep])
}

View File

@ -7,12 +7,16 @@ import (
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
)
func (a *Agent) buildMemoryContext(input string) string {
func (a *Agent) buildMemoryContext(input string, maxTokens int) string {
if a.indexer == nil {
return ""
}
injected := a.indexer.BuildContext(input)
return a.indexer.FormatContext(injected)
s := a.indexer.FormatContext(injected)
if maxTokens > 0 {
s = TruncateByTokens(s, maxTokens)
}
return s
}
func (a *Agent) buildSystemPrompt(memContext string, userInput string) string {