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