mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-03 07:43:58 +00:00
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:
@ -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),
|
||||
|
||||
}
|
||||
|
||||
@ -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 == "" {
|
||||
|
||||
@ -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) {
|
||||
|
||||
@ -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)
|
||||
}
|
||||
|
||||
@ -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})
|
||||
|
||||
86
internal/agent/core/tokenbudget.go
Normal file
86
internal/agent/core/tokenbudget.go
Normal 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])
|
||||
}
|
||||
@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user