Files
ai-gateway-go/internal/workbench/retrievers_test.go
T
superidou b536672000 feat(m8): P2 pgvector + Ollama 向量化与语义检索
- PostgreSQL 切换 pgvector/pgvector:pg17 镜像;迁移 000024 建 vector 扩展、
  knowledge_chunks.embedding vector(1024) + HNSW 余弦索引,retrieval_mode 放宽三态
- OllamaEmbedder 本地 bge-m3 批量嵌入,404 惰性 pull 重试,维度/超时校验,可整体关闭
- SemanticRetriever/HybridRetriever + NewRetriever 按 retrieval_mode 分发,缺 embedder 回退 FTS
- 文档入库同步批量向量化;Ollama 故障降级入库 + embedding_failed 事件
- 修复 pgx CopyFrom 对 vector 列二进制编码误读:COPY 基础列后同事务 unnest 批量回填
- 修复降级路径 embeddings=nil 索引越界 panic(Add 与 Reprocess)
- 知识库列表 vectorized_chunk_count + 前端三态检索模式选择与向量化覆盖率
- 单测 embedder/retrievers + 集成 TestKnowledgeVectorLifecycle 全绿

Co-Authored-By: Claude <noreply@anthropic.com>
2026-08-12 15:16:32 +08:00

109 lines
4.1 KiB
Go

package workbench
import (
"context"
"errors"
"testing"
)
func TestValidateKnowledgeBaseRetrievalModes(t *testing.T) {
for _, mode := range []string{"postgres_fts", "vector", "hybrid"} {
kb := KnowledgeBase{Name: "kb", RetrievalMode: mode, ChunkSize: 800, ChunkOverlap: 100}
if err := validateKnowledgeBase(&kb); err != nil {
t.Fatalf("mode %s should be accepted, got %v", mode, err)
}
}
kb := KnowledgeBase{Name: "kb", RetrievalMode: "bm25", ChunkSize: 800, ChunkOverlap: 100}
if err := validateKnowledgeBase(&kb); err == nil {
t.Fatal("unknown retrieval mode should be rejected")
}
// 空模式回填默认值
empty := KnowledgeBase{Name: "kb", ChunkSize: 800, ChunkOverlap: 100}
if err := validateKnowledgeBase(&empty); err != nil || empty.RetrievalMode != "postgres_fts" {
t.Fatalf("empty mode should default to postgres_fts, got %q err=%v", empty.RetrievalMode, err)
}
}
func TestValidateSearchClamps(t *testing.T) {
if _, topK, err := validateSearch(" 关键词 ", 0); err != nil || topK != 4 {
t.Fatalf("empty topK should clamp to 4, got %d err=%v", topK, err)
}
if _, topK, err := validateSearch("关键词", 99); err != nil || topK != 20 {
t.Fatalf("topK should clamp to 20, got %d err=%v", topK, err)
}
if _, _, err := validateSearch(" ", 4); err == nil {
t.Fatal("blank query should be rejected")
}
long := make([]byte, 6001)
for i := range long {
long[i] = 'a'
}
if _, _, err := validateSearch(string(long), 4); err == nil {
t.Fatal("over-long query should be rejected")
}
}
// failingEmbedder 总是报错,用于验证 computeEmbeddings 优雅降级。
type failingEmbedder struct{}
func (failingEmbedder) Embed(context.Context, []string) ([][]float32, error) { return nil, errors.New("embedding service down") }
func (failingEmbedder) Dim() int { return 1024 }
func TestComputeEmbeddingsDispatch(t *testing.T) {
service := NewService(nil)
chunks := []string{"a", "b"}
// 未注入 embedder:任何模式都不算向量。
if result := service.computeEmbeddings(context.Background(), KnowledgeBase{ID: "kb", RetrievalMode: "vector"}, chunks); result.needed {
t.Fatal("no embedder should not mark embedding needed")
}
// 已注入 embedder 但模式为 postgres_fts:不算向量。
service.SetEmbedder(fakeGoodEmbedder{})
if result := service.computeEmbeddings(context.Background(), KnowledgeBase{ID: "kb", RetrievalMode: "postgres_fts"}, chunks); result.needed {
t.Fatal("postgres_fts mode should not embed")
}
// vector 模式 + 失败 embedder:needed + degraded,embedding 为 nil。
service.SetEmbedder(failingEmbedder{})
result := service.computeEmbeddings(context.Background(), KnowledgeBase{ID: "kb", RetrievalMode: "hybrid"}, chunks)
if !result.needed || !result.degraded || result.embeddings != nil {
t.Fatalf("expected needed+degraded with nil embeddings, got %+v", result)
}
// vector 模式 + 正常 embedder:返回与 chunks 等长的向量。
service.SetEmbedder(fakeGoodEmbedder{})
result = service.computeEmbeddings(context.Background(), KnowledgeBase{ID: "kb", RetrievalMode: "vector"}, chunks)
if !result.needed || result.degraded || len(result.embeddings) != len(chunks) {
t.Fatalf("expected aligned embeddings, got %+v", result)
}
}
type fakeGoodEmbedder struct{}
func (fakeGoodEmbedder) Embed(_ context.Context, texts []string) ([][]float32, error) {
vectors := make([][]float32, len(texts))
for i := range texts {
vectors[i] = make([]float32, 4)
for j := range vectors[i] {
vectors[i][j] = float32(i + 1)
}
}
return vectors, nil
}
func (fakeGoodEmbedder) Dim() int { return 4 }
func TestNewRetrieverNilEmbedder(t *testing.T) {
// embedder 为 nil 时 NewRetriever 返回的检索器直接走 FTS 分支,不会查询检索模式。
service := NewService(nil)
retriever := NewRetriever(service, nil)
if retriever == nil {
t.Fatal("NewRetriever with nil embedder should not return nil")
}
// FTS 分支在无连接池时只在 query 校验处返回,不 panic。
if _, err := retriever.Search(context.Background(), "kb", " ", 4); err == nil {
t.Fatal("blank query should fail at validation before touching the pool")
}
}