b536672000
- 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>
109 lines
4.1 KiB
Go
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")
|
|
}
|
|
}
|