mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-03 07:43:58 +00:00
feat: multi-language embedding with CleanTemplateText + three-branch vector strategy
- StaticEmbedder: pre-trained ConceptNet Numberbatch/fastText word embeddings, auto-download with TF-IDF fallback, comma-separated multi-model paths - CleanTemplateText: regex stripping of QQ tool call templates and noise - textForVector: per-source vector strategy (agent→Response, user→Input, cold_storage→both) - Indexer.BuildContext and ExtractKeywords now clean input before vectorization - Protect recent 10 events in Prune (regression fix: use local var not const)
This commit is contained in:
@ -1,12 +1,19 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||
)
|
||||
|
||||
func newTestCtx() *RelevanceContext {
|
||||
return NewRelevanceContext("", memory.NewStaticEmbedder(""))
|
||||
}
|
||||
|
||||
func TestContextAppendAndLen(t *testing.T) {
|
||||
ctx := NewRelevanceContext("")
|
||||
ctx := newTestCtx()
|
||||
if ctx.Len() != 0 {
|
||||
t.Errorf("new context should be empty, got %d", ctx.Len())
|
||||
}
|
||||
@ -18,7 +25,7 @@ func TestContextAppendAndLen(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestContextRecent(t *testing.T) {
|
||||
ctx := NewRelevanceContext("")
|
||||
ctx := newTestCtx()
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "a"})
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "b"})
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "c"})
|
||||
@ -33,7 +40,7 @@ func TestContextRecent(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestContextFormat(t *testing.T) {
|
||||
ctx := NewRelevanceContext("")
|
||||
ctx := newTestCtx()
|
||||
f := ctx.Format()
|
||||
if f != "" {
|
||||
t.Errorf("empty context should format to empty string, got %q", f)
|
||||
@ -54,7 +61,7 @@ func TestContextFormat(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestContextPruneKeepsTopK(t *testing.T) {
|
||||
ctx := NewRelevanceContext("")
|
||||
ctx := newTestCtx()
|
||||
for i := 0; i < 20; i++ {
|
||||
ctx.Append(ContextEvent{
|
||||
Timestamp: time.Now(),
|
||||
@ -63,7 +70,6 @@ func TestContextPruneKeepsTopK(t *testing.T) {
|
||||
Response: "是的天气不错",
|
||||
})
|
||||
}
|
||||
// 加一条不同主题的
|
||||
ctx.Append(ContextEvent{
|
||||
Timestamp: time.Now(),
|
||||
Source: "user",
|
||||
@ -71,17 +77,16 @@ func TestContextPruneKeepsTopK(t *testing.T) {
|
||||
Response: "好的我来算",
|
||||
})
|
||||
|
||||
archived := ctx.Prune("微积分", 5, nil) // nil docStore → 不归档,只裁剪
|
||||
archived := ctx.Prune("微积分", 5, nil)
|
||||
_ = archived
|
||||
|
||||
// protectCount=10 + topK=5 → 最多保留 15
|
||||
if ctx.Len() > 15 {
|
||||
t.Errorf("after prune to 5, len should be ≤15, got %d", ctx.Len())
|
||||
}
|
||||
}
|
||||
|
||||
func TestContextPruneWithDocStore(t *testing.T) {
|
||||
ctx := NewRelevanceContext("")
|
||||
ctx := newTestCtx()
|
||||
for i := 0; i < 15; i++ {
|
||||
ctx.Append(ContextEvent{
|
||||
Timestamp: time.Now(),
|
||||
@ -98,8 +103,7 @@ func TestContextPruneWithDocStore(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestContextAppendAfterPrune(t *testing.T) {
|
||||
ctx := NewRelevanceContext("")
|
||||
// 需要超过 protectCount(10) + topK(3) 个事件才能产生修剪候选
|
||||
ctx := newTestCtx()
|
||||
for i := 0; i < 20; i++ {
|
||||
ctx.Append(ContextEvent{
|
||||
Timestamp: time.Now(),
|
||||
@ -109,7 +113,7 @@ func TestContextAppendAfterPrune(t *testing.T) {
|
||||
}
|
||||
|
||||
ctx.Prune("hello", 3, nil)
|
||||
if ctx.Len() > 13 { // 10 protected + 3 topK
|
||||
if ctx.Len() > 13 {
|
||||
t.Errorf("expected ≤13 after prune, got %d", ctx.Len())
|
||||
}
|
||||
|
||||
@ -119,6 +123,107 @@ func TestContextAppendAfterPrune(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestContextPruneWithStaticEmbedder(t *testing.T) {
|
||||
tmpFile, err := os.CreateTemp("", "test_embeddings_*.txt")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer os.Remove(tmpFile.Name())
|
||||
|
||||
content := `8 4
|
||||
天气 0.1 0.2 0.3 0.4
|
||||
下雨 0.15 0.25 0.35 0.45
|
||||
台风 0.12 0.22 0.32 0.42
|
||||
股票 0.9 0.1 0.1 0.1
|
||||
基金 0.85 0.15 0.1 0.1
|
||||
微积分 0.1 0.1 0.9 0.1
|
||||
导数 0.15 0.1 0.85 0.15
|
||||
数学 0.1 0.1 0.8 0.2
|
||||
`
|
||||
if _, err := tmpFile.WriteString(content); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tmpFile.Close()
|
||||
|
||||
embedder := memory.NewStaticEmbedder(tmpFile.Name())
|
||||
if !embedder.Loaded() {
|
||||
t.Fatal("embedder should be loaded")
|
||||
}
|
||||
|
||||
ctx := NewRelevanceContext("", embedder)
|
||||
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "今天天气很好", Response: "是的"})
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "会不会下雨", Response: "会"})
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "台风来了", Response: "注意"})
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "帮我算微积分", Response: "好的"})
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "导数怎么求", Response: "公式"})
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "数学题", Response: "解答"})
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "股票涨了", Response: "恭喜"})
|
||||
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "基金定投", Response: "可以"})
|
||||
|
||||
if ctx.Len() != 8 {
|
||||
t.Fatalf("expected 8 events, got %d", ctx.Len())
|
||||
}
|
||||
|
||||
archived := ctx.Prune("最近基金怎么样", 3, nil)
|
||||
|
||||
if ctx.Len() > 13 {
|
||||
t.Errorf("prune should limit total events, got %d", ctx.Len())
|
||||
}
|
||||
|
||||
remaining := ctx.Format()
|
||||
t.Logf("query: 最近基金怎么样\nremaining events:\n%s", remaining)
|
||||
t.Logf("archived: %d", archived)
|
||||
|
||||
needsFund := contains(remaining, "基金定投") || contains(remaining, "股票涨了")
|
||||
needsWeather := contains(remaining, "今天天气很好") || contains(remaining, "台风来了")
|
||||
|
||||
t.Logf("has financial events: %v, has weather events: %v", needsFund, needsWeather)
|
||||
}
|
||||
|
||||
func TestContextPruneRecent10Protected(t *testing.T) {
|
||||
ctx := newTestCtx()
|
||||
|
||||
for i := 0; i < 15; i++ {
|
||||
ctx.Append(ContextEvent{
|
||||
Timestamp: time.Now(),
|
||||
Source: "user",
|
||||
Input: "今天天气很好",
|
||||
})
|
||||
}
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
ctx.Append(ContextEvent{
|
||||
Timestamp: time.Now(),
|
||||
Source: "user",
|
||||
Input: "股票行情",
|
||||
})
|
||||
}
|
||||
|
||||
ctx.Prune("天气", 3, nil)
|
||||
|
||||
// 最近 10 条全部是"股票行情"(第6-15条是天气,第16-20条是股票)
|
||||
// protectCount=10 保护最近 10 条 → 5 条天气最多保留 5+3=8 条
|
||||
// 至少最近 10 条全部保留 → 至少包含 5 条股票
|
||||
remaining := ctx.Format()
|
||||
t.Logf("after weather query:\n%s", remaining)
|
||||
weatherCount := 0
|
||||
stockCount := 0
|
||||
for _, line := range splitLines(remaining) {
|
||||
if contains(line, "天气") {
|
||||
weatherCount++
|
||||
}
|
||||
if contains(line, "股票") {
|
||||
stockCount++
|
||||
}
|
||||
}
|
||||
t.Logf("weather events: %d, stock events: %d", weatherCount, stockCount)
|
||||
|
||||
if stockCount < 5 {
|
||||
t.Errorf("recent 10 should all be protected, expected at least 5 stock events, got %d", stockCount)
|
||||
}
|
||||
}
|
||||
|
||||
func contains(s, substr string) bool {
|
||||
return len(s) >= len(substr) && containsStr(s, substr)
|
||||
}
|
||||
@ -131,3 +236,18 @@ func containsStr(s, substr string) bool {
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func splitLines(s string) []string {
|
||||
var lines []string
|
||||
start := 0
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i] == '\n' {
|
||||
lines = append(lines, s[start:i])
|
||||
start = i + 1
|
||||
}
|
||||
}
|
||||
if start < len(s) {
|
||||
lines = append(lines, s[start:])
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user