Files
HomeAgent/internal/memory/media/media_vec_test.go
JianFeeeee 8c9d96e065 feat(media): media.Store 加向量存储与跨模态检索
media.Item 新增 Vec []float64 和 VecModel 字段,视觉嵌入向量以 JSON
TEXT 存库(为什么不存 BLOB:Go 的 json.Marshal 对 []float64 是自然的,
而 SQLite BLOB 是 []byte 多一层序列化;单条最多 23KB,TEXT 够用)。

initSchema 加 ALTER TABLE 迁移 vec/vec_model 两列(幂等)。scanItem
扩展读回。新增 SetVec 和 QueryMedia 方法。

QueryMedia 对所有已嵌入媒体做余弦相似度检索,维度不一致的项自动跳过
——这是跨模态检索的核心:查询可以是图片也可以是文本,被查的媒体库
里每个 item 也有视觉向量,两者在同一空间比对,谁的相似度更高就召回谁。

配套 5 个单测覆盖:基本相似度排序、无向量项被跳过、维度不匹配过滤、
空查询安全、向量持久化正确性。
2026-09-07 22:31:39 +08:00

136 lines
3.9 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package media
import (
"math"
"testing"
)
func TestQueryMedia_BasicSimilarity(t *testing.T) {
s := newTestStore(t, 0)
defer s.Close()
// 入库三张带向量的媒体:两张图、一段音频
d1, _ := s.Put([]byte("img1"), Item{MIME: "image/png", Description: "紫蓝红三色带"})
d2, _ := s.Put([]byte("img2"), Item{MIME: "image/jpeg", Description: "蓝紫红渐变"})
d3, _ := s.Put([]byte("aud1"), Item{MIME: "audio/wav", Description: "一段语音"})
// 模拟视觉嵌入img1 和 img2 向量接近aud1 远离
vec1 := []float64{0.9, 0.1, 0.0, 0.0}
vec2 := []float64{0.8, 0.2, 0.0, 0.0} // 与 vec1 相似
vec3 := []float64{0.0, 0.0, 0.9, 0.1} // 与前两个完全不同
s.SetVec(d1, vec1, "test-clip")
s.SetVec(d2, vec2, "test-clip")
s.SetVec(d3, vec3, "test-clip")
// 用 vec1 作为查询vec2 最相似vec3 与 vec1 正交(相似度 0被阈值过滤
results, err := s.QueryMedia(vec1, "test-clip", 10)
if err != nil {
t.Fatal(err)
}
// vec3 与 vec1 正交(余弦相似度 0被 0.05 阈值正确剔除 → 只召回 2 个
if len(results) != 2 {
t.Fatalf("expected 2 results (正交的 aud1 被阈值过滤), got %d", len(results))
}
// 第一个应该是 img20.9 vs d1 的 1.0?不,这里算清楚)
// vec1·vec2 与 vec1·vec1 比较:
// sim(vec1,vec1) = 1.0img1 与自身sim(vec1,vec2) = 0.9*0.8+0.1*0.2 = 0.74
// 所以 img1自相似 1.0排第一img2 排第二
if results[0].Digest != d1 {
t.Errorf("expected d1 (自相似 1.0) as first, got %s", results[0].Digest)
}
if results[1].Digest != d2 {
t.Errorf("expected d2 as second, got %s", results[1].Digest)
}
// 验证分数img1 与自身是 1.0
selfScore := cosineSimilaritySlice(vec1, vec1)
if math.Abs(selfScore-1.0) > 1e-10 {
t.Errorf("self-similarity should be 1.0, got %f", selfScore)
}
// img1 与 aud1 的相似度应该很低
crossScore := cosineSimilaritySlice(vec1, vec3)
if crossScore > 0.1 {
t.Errorf("cross-modality similarity should be low, got %f", crossScore)
}
}
func TestQueryMedia_EmptyVecSkipped(t *testing.T) {
s := newTestStore(t, 0)
defer s.Close()
d1, _ := s.Put([]byte("img1"), Item{MIME: "image/png"})
_, _ = s.Put([]byte("img2"), Item{MIME: "image/png"})
// d1 有向量d2 没有
s.SetVec(d1, []float64{0.5, 0.5}, "test")
// d2 留空
results, err := s.QueryMedia([]float64{0.5, 0.5}, "test", 10)
if err != nil {
t.Fatal(err)
}
if len(results) != 1 {
t.Fatalf("expected 1 result (d2 has no vec), got %d", len(results))
}
if results[0].Digest != d1 {
t.Errorf("expected d1, got %s", results[0].Digest)
}
}
func TestQueryMedia_DimensionMismatchSkipped(t *testing.T) {
s := newTestStore(t, 0)
defer s.Close()
d1, _ := s.Put([]byte("img1"), Item{MIME: "image/png"})
s.SetVec(d1, []float64{0.5, 0.5}, "model-A") // 2 维
// 查询用 3 维向量:维度不匹配,应该返回空
results, err := s.QueryMedia([]float64{0.3, 0.3, 0.3}, "model-A", 10)
if err != nil {
t.Fatal(err)
}
if len(results) != 0 {
t.Fatalf("expected 0 results (dim mismatch), got %d", len(results))
}
}
func TestQueryMedia_EmptyQueryReturnsNil(t *testing.T) {
s := newTestStore(t, 0)
defer s.Close()
results, err := s.QueryMedia(nil, "", 10)
if err != nil {
t.Fatal(err)
}
if results != nil {
t.Fatalf("expected nil, got %d results", len(results))
}
}
func TestSetVec_PersistsCorrectly(t *testing.T) {
s := newTestStore(t, 0)
defer s.Close()
d, _ := s.Put([]byte("hello"), Item{MIME: "image/png"})
vec := []float64{0.1, 0.2, 0.3, 0.4}
s.SetVec(d, vec, "clip-vit-b32")
it, err := s.Stat(d)
if err != nil {
t.Fatal(err)
}
if it.VecModel != "clip-vit-b32" {
t.Errorf("VecModel = %q, want clip-vit-b32", it.VecModel)
}
if len(it.Vec) != 4 {
t.Fatalf("Vec len = %d, want 4", len(it.Vec))
}
for i, v := range vec {
if math.Abs(it.Vec[i]-v) > 1e-10 {
t.Errorf("Vec[%d] = %f, want %f", i, it.Vec[i], v)
}
}
}