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>
160 lines
4.9 KiB
Go
160 lines
4.9 KiB
Go
package workbench
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"sort"
|
|
"strconv"
|
|
"strings"
|
|
|
|
"github.com/jackc/pgx/v5"
|
|
)
|
|
|
|
// NewRetriever 按知识库的 retrieval_mode 分发检索器:
|
|
// - embedder 为 nil 时恒为 postgres_fts(向量化未启用)
|
|
// - 否则按 knowledge_bases.retrieval_mode 每次 Search 动态分发
|
|
// (postgres_fts / vector / hybrid),管理员修改模式后无需重启即生效。
|
|
func NewRetriever(service *Service, embedder Embedder) Retriever {
|
|
return &modeRetriever{pool: service.pool, embedder: embedder}
|
|
}
|
|
|
|
type modeRetriever struct {
|
|
pool interface {
|
|
Query(context.Context, string, ...any) (pgx.Rows, error)
|
|
QueryRow(context.Context, string, ...any) pgx.Row
|
|
}
|
|
embedder Embedder
|
|
}
|
|
|
|
func (r *modeRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]SearchHit, error) {
|
|
if r.embedder == nil {
|
|
return (&PostgreSQLRetriever{pool: r.pool}).Search(ctx, knowledgeBaseID, query, topK)
|
|
}
|
|
var mode string
|
|
err := r.pool.QueryRow(ctx, `SELECT retrieval_mode FROM gateway.knowledge_bases WHERE id=$1`, knowledgeBaseID).Scan(&mode)
|
|
if err != nil {
|
|
return nil, mapNotFound(err)
|
|
}
|
|
switch mode {
|
|
case "vector":
|
|
return (&SemanticRetriever{pool: r.pool, embedder: r.embedder}).Search(ctx, knowledgeBaseID, query, topK)
|
|
case "hybrid":
|
|
return (&HybridRetriever{pool: r.pool, embedder: r.embedder}).Search(ctx, knowledgeBaseID, query, topK)
|
|
default:
|
|
return (&PostgreSQLRetriever{pool: r.pool}).Search(ctx, knowledgeBaseID, query, topK)
|
|
}
|
|
}
|
|
|
|
// SemanticRetriever 用 pgvector 余弦距离召回最近分块。
|
|
type SemanticRetriever struct {
|
|
pool interface {
|
|
Query(context.Context, string, ...any) (pgx.Rows, error)
|
|
}
|
|
embedder Embedder
|
|
}
|
|
|
|
func (r *SemanticRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]SearchHit, error) {
|
|
query, topK, err := validateSearch(query, topK)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
vectors, err := r.embedder.Embed(ctx, []string{query})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if len(vectors) == 0 {
|
|
return nil, errors.New("查询向量为空")
|
|
}
|
|
rows, err := r.pool.Query(ctx, `SELECT c.id::text,c.document_id::text,d.title,c.chunk_index,c.content,
|
|
(1-(c.embedding <=> $2::vector))::float8 AS score
|
|
FROM gateway.knowledge_chunks c JOIN gateway.knowledge_documents d ON d.id=c.document_id JOIN gateway.knowledge_bases k ON k.id=c.knowledge_base_id
|
|
WHERE c.knowledge_base_id=$1 AND k.enabled AND d.status='ready' AND c.embedding IS NOT NULL
|
|
ORDER BY c.embedding <=> $2::vector LIMIT $3`, knowledgeBaseID, formatVector(vectors[0]), topK)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
hits := []SearchHit{}
|
|
for rows.Next() {
|
|
var h SearchHit
|
|
if err = rows.Scan(&h.ChunkID, &h.DocumentID, &h.DocumentTitle, &h.ChunkIndex, &h.Content, &h.Score); err != nil {
|
|
return nil, err
|
|
}
|
|
hits = append(hits, h)
|
|
}
|
|
return hits, rows.Err()
|
|
}
|
|
|
|
// HybridRetriever 融合 FTS 与语义召回:分别取 topK*2,按 chunk_id 去重取高分,
|
|
// 再截断到 topK。任一子检索失败时回退到另一路,保证向量化降级仍有结果。
|
|
type HybridRetriever struct {
|
|
pool interface {
|
|
Query(context.Context, string, ...any) (pgx.Rows, error)
|
|
}
|
|
embedder Embedder
|
|
}
|
|
|
|
func (r *HybridRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]SearchHit, error) {
|
|
_, topK, err := validateSearch(query, topK)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
fts, ftsErr := (&PostgreSQLRetriever{pool: r.pool}).Search(ctx, knowledgeBaseID, query, topK*2)
|
|
sem, semErr := (&SemanticRetriever{pool: r.pool, embedder: r.embedder}).Search(ctx, knowledgeBaseID, query, topK*2)
|
|
if ftsErr != nil && semErr != nil {
|
|
return nil, ftsErr
|
|
}
|
|
best := make(map[string]SearchHit)
|
|
for _, h := range fts {
|
|
if existing, ok := best[h.ChunkID]; !ok || h.Score > existing.Score {
|
|
best[h.ChunkID] = h
|
|
}
|
|
}
|
|
for _, h := range sem {
|
|
if existing, ok := best[h.ChunkID]; !ok || h.Score > existing.Score {
|
|
best[h.ChunkID] = h
|
|
}
|
|
}
|
|
merged := make([]SearchHit, 0, len(best))
|
|
for _, h := range best {
|
|
merged = append(merged, h)
|
|
}
|
|
sort.Slice(merged, func(i, j int) bool { return merged[i].Score > merged[j].Score })
|
|
if len(merged) > topK {
|
|
merged = merged[:topK]
|
|
}
|
|
return merged, nil
|
|
}
|
|
|
|
func validateSearch(query string, topK int) (string, int, error) {
|
|
query = strings.TrimSpace(query)
|
|
if query == "" {
|
|
return "", 0, errors.New("检索词不能为空")
|
|
}
|
|
if len(query) > 6000 {
|
|
return "", 0, errors.New("检索词过长")
|
|
}
|
|
if topK < 1 {
|
|
topK = 4
|
|
}
|
|
if topK > 20 {
|
|
topK = 20
|
|
}
|
|
return query, topK, nil
|
|
}
|
|
|
|
// formatVector 把 []float32 转成 pgvector 字面量字符串 "[0.1,0.2,...]",
|
|
// 直接以 $n::vector 参数传入,避免引入 pgvector-go 依赖。
|
|
func formatVector(vector []float32) string {
|
|
var builder strings.Builder
|
|
builder.WriteByte('[')
|
|
for index, value := range vector {
|
|
if index > 0 {
|
|
builder.WriteByte(',')
|
|
}
|
|
builder.WriteString(strconv.FormatFloat(float64(value), 'g', -1, 32))
|
|
}
|
|
builder.WriteByte(']')
|
|
return builder.String()
|
|
}
|