Files
ai-gateway-go/internal/workbench/retrievers.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

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()
}