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

501 lines
18 KiB
Go

package workbench
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"log/slog"
"regexp"
"sort"
"strings"
"unicode"
"unicode/utf8"
"github.com/jackc/pgx/v5"
)
type Retriever interface {
Search(context.Context, string, string, int) ([]SearchHit, error)
}
type PostgreSQLRetriever struct {
pool interface {
Query(context.Context, string, ...any) (pgx.Rows, error)
}
}
func (r *PostgreSQLRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]SearchHit, error) {
var err error
query, topK, err = validateSearch(query, topK)
if err != nil {
return nil, err
}
tokens := searchTokens(query)
rows, err := r.pool.Query(ctx, `WITH q AS (SELECT plainto_tsquery('simple',$2) AS tsq,lower($2) AS raw), tokens AS (SELECT unnest($4::text[]) AS token)
SELECT c.id::text,c.document_id::text,d.title,c.chunk_index,c.content,
greatest(ts_rank_cd(c.search_vector,q.tsq),CASE WHEN strpos(lower(c.content),q.raw)>0 THEN 1.0 ELSE 0.0 END,
(SELECT count(*)::float8/greatest(array_length($4::text[],1),1) FROM tokens WHERE strpos(lower(c.content),token)>0))::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 CROSS JOIN q
WHERE c.knowledge_base_id=$1 AND k.enabled AND d.status='ready' AND (c.search_vector @@ q.tsq OR strpos(lower(c.content),q.raw)>0 OR EXISTS(SELECT 1 FROM tokens WHERE strpos(lower(c.content),token)>0))
ORDER BY score DESC,c.document_id,c.chunk_index LIMIT $3`, knowledgeBaseID, query, topK, tokens)
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()
}
func searchTokens(query string) []string {
query = strings.ToLower(strings.TrimSpace(query))
seen := map[string]struct{}{}
tokens := make([]string, 0, 32)
add := func(value string) {
value = strings.TrimSpace(value)
if utf8.RuneCountInString(value) < 2 {
return
}
if _, ok := seen[value]; ok {
return
}
seen[value] = struct{}{}
tokens = append(tokens, value)
}
words := strings.FieldsFunc(query, func(value rune) bool { return !unicode.IsLetter(value) && !unicode.IsNumber(value) })
for _, word := range words {
runes := []rune(word)
hasCJK := false
for _, value := range runes {
if unicode.In(value, unicode.Han) {
hasCJK = true
break
}
}
if !hasCJK {
add(word)
continue
}
for index := 0; index+1 < len(runes); index++ {
add(string(runes[index : index+2]))
}
}
if len(tokens) > 64 {
tokens = tokens[:64]
}
sort.Strings(tokens)
return tokens
}
func validateKnowledgeBase(k *KnowledgeBase) error {
k.Name = strings.TrimSpace(k.Name)
k.Description = strings.TrimSpace(k.Description)
if k.Name == "" || len(k.Name) > 128 || len(k.Description) > 4000 {
return errors.New("知识库名称或描述格式无效")
}
if k.RetrievalMode == "" {
k.RetrievalMode = "postgres_fts"
}
switch k.RetrievalMode {
case "postgres_fts", "vector", "hybrid":
// M8 P2:三态检索模式。vector/hybrid 需要向量化器,未启用时检索自动回退 FTS。
default:
return errors.New("retrieval_mode 仅支持 postgres_fts / vector / hybrid")
}
if k.ChunkSize == 0 {
k.ChunkSize = 800
}
if k.ChunkSize < 200 || k.ChunkSize > 8000 {
return errors.New("chunk_size 应在 200-8000 之间")
}
if k.ChunkOverlap < 0 || k.ChunkOverlap > k.ChunkSize/2 {
return errors.New("chunk_overlap 应在 0 到 chunk_size 一半之间")
}
var err error
k.DepartmentIDs, err = normalizeStrings(k.DepartmentIDs, 100)
return err
}
const knowledgeBaseSelect = `SELECT k.id::text,k.name,k.description,k.retrieval_mode,k.chunk_size,k.chunk_overlap,k.department_ids::text[],k.enabled,k.revision,
count(DISTINCT d.id)::int,count(c.id)::int,count(c.embedding)::int,k.created_at,k.updated_at
FROM gateway.knowledge_bases k LEFT JOIN gateway.knowledge_documents d ON d.knowledge_base_id=k.id LEFT JOIN gateway.knowledge_chunks c ON c.document_id=d.id`
func scanKnowledgeBase(row pgx.Row) (KnowledgeBase, error) {
var k KnowledgeBase
err := row.Scan(&k.ID, &k.Name, &k.Description, &k.RetrievalMode, &k.ChunkSize, &k.ChunkOverlap, &k.DepartmentIDs, &k.Enabled, &k.Revision, &k.DocumentCount, &k.ChunkCount, &k.VectorizedChunkCount, &k.CreatedAt, &k.UpdatedAt)
return k, mapNotFound(err)
}
func (s *Service) ListKnowledgeBases(ctx context.Context) ([]KnowledgeBase, error) {
rows, err := s.pool.Query(ctx, knowledgeBaseSelect+` GROUP BY k.id ORDER BY k.updated_at DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
items := []KnowledgeBase{}
for rows.Next() {
k, err := scanKnowledgeBase(rows)
if err != nil {
return nil, err
}
items = append(items, k)
}
return items, rows.Err()
}
func (s *Service) GetKnowledgeBase(ctx context.Context, id string) (KnowledgeBase, error) {
return scanKnowledgeBase(s.pool.QueryRow(ctx, knowledgeBaseSelect+` WHERE k.id=$1 GROUP BY k.id`, id))
}
func (s *Service) SaveKnowledgeBase(ctx context.Context, k KnowledgeBase, actorID string, create bool) (KnowledgeBase, error) {
if err := validateKnowledgeBase(&k); err != nil {
return KnowledgeBase{}, err
}
tx, err := s.pool.Begin(ctx)
if err != nil {
return KnowledgeBase{}, err
}
defer rollback(ctx, tx)
if create {
k.ID, err = newUUID()
if err != nil {
return KnowledgeBase{}, err
}
_, err = tx.Exec(ctx, `INSERT INTO gateway.knowledge_bases(id,name,description,retrieval_mode,chunk_size,chunk_overlap,department_ids,enabled,created_by) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9)`, k.ID, k.Name, k.Description, k.RetrievalMode, k.ChunkSize, k.ChunkOverlap, k.DepartmentIDs, k.Enabled, actorID)
} else {
tag, updateErr := tx.Exec(ctx, `UPDATE gateway.knowledge_bases SET name=$2,description=$3,retrieval_mode=$4,chunk_size=$5,chunk_overlap=$6,department_ids=$7,enabled=$8,revision=revision+1,updated_at=clock_timestamp() WHERE id=$1`, k.ID, k.Name, k.Description, k.RetrievalMode, k.ChunkSize, k.ChunkOverlap, k.DepartmentIDs, k.Enabled)
err = updateErr
if err == nil && tag.RowsAffected() == 0 {
return KnowledgeBase{}, ErrNotFound
}
}
if err != nil {
return KnowledgeBase{}, err
}
event := "knowledge_base.updated"
if create {
event = "knowledge_base.created"
}
if err = emit(ctx, tx, event, "knowledge_base", k.ID, actorID, nil); err != nil {
return KnowledgeBase{}, err
}
if err = tx.Commit(ctx); err != nil {
return KnowledgeBase{}, err
}
return s.GetKnowledgeBase(ctx, k.ID)
}
func (s *Service) DeleteKnowledgeBase(ctx context.Context, id, actorID string) error {
tx, err := s.pool.Begin(ctx)
if err != nil {
return err
}
defer rollback(ctx, tx)
var used bool
if err = tx.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM gateway.applications WHERE draft_config->'knowledge_base_ids' ? $1 UNION ALL SELECT 1 FROM gateway.application_versions WHERE config->'knowledge_base_ids' ? $1)`, id).Scan(&used); err != nil {
return err
}
if used {
return ErrConflict
}
tag, err := tx.Exec(ctx, `DELETE FROM gateway.knowledge_bases WHERE id=$1`, id)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrNotFound
}
if err = emit(ctx, tx, "knowledge_base.deleted", "knowledge_base", id, actorID, nil); err != nil {
return err
}
return tx.Commit(ctx)
}
func ChunkText(text string, size, overlap int) []string {
if size < 200 {
size = 800
}
if overlap < 0 {
overlap = 0
}
if overlap > size/2 {
overlap = size / 2
}
paragraphs := regexp.MustCompile(`\n\s*\n`).Split(strings.TrimSpace(text), -1)
chunks := []string{}
current := []rune{}
appendCurrent := func() {
value := strings.TrimSpace(string(current))
if value != "" {
chunks = append(chunks, value)
}
current = nil
}
for _, paragraph := range paragraphs {
runes := []rune(strings.TrimSpace(paragraph))
if len(runes) == 0 {
continue
}
if len(runes) > size {
appendCurrent()
step := size - overlap
if step < 1 {
step = 1
}
for start := 0; start < len(runes); start += step {
end := start + size
if end > len(runes) {
end = len(runes)
}
chunks = append(chunks, string(runes[start:end]))
if end == len(runes) {
break
}
}
continue
}
separator := 0
if len(current) > 0 {
separator = 2
}
if len(current)+separator+len(runes) <= size {
if separator > 0 {
current = append(current, '\n', '\n')
}
current = append(current, runes...)
continue
}
previous := append([]rune(nil), current...)
appendCurrent()
if overlap > 0 && len(previous) > 0 {
start := len(previous) - overlap
if start < 0 {
start = 0
}
current = append(current, previous[start:]...)
current = append(current, '\n', '\n')
}
current = append(current, runes...)
}
appendCurrent()
return chunks
}
// embedResult 描述文档入库时的向量化结果。
type embedResult struct {
embeddings [][]float32 // 与 chunks 一一对应;needed=false 或 degraded=true 时为 nil
needed bool // 该知识库模式需要 embedding
degraded bool // 需要但计算失败,文档降级为纯 FTS 入库
}
// computeEmbeddings 在 KB 为 vector/hybrid 且已注入 embedder 时同步批量计算分块向量。
// 失败时不阻断入库:返回 degraded=true,由调用方在事务内补发 embedding_failed 事件。
func (s *Service) computeEmbeddings(ctx context.Context, kb KnowledgeBase, chunks []string) embedResult {
if s.embedder == nil || (kb.RetrievalMode != "vector" && kb.RetrievalMode != "hybrid") {
return embedResult{}
}
vectors, err := s.embedder.Embed(ctx, chunks)
if err != nil {
slog.Warn("knowledge embedding failed, storing document without vectors", "knowledge_base_id", kb.ID, "error", err)
return embedResult{needed: true, degraded: true}
}
if len(vectors) != len(chunks) {
slog.Warn("knowledge embedding count mismatch, storing document without vectors", "knowledge_base_id", kb.ID, "got", len(vectors), "want", len(chunks))
return embedResult{needed: true, degraded: true}
}
return embedResult{embeddings: vectors, needed: true}
}
func (s *Service) AddKnowledgeDocument(ctx context.Context, kbID, title, sourceType, sourceURI, content, actorID string) (KnowledgeDocument, error) {
title = strings.TrimSpace(title)
sourceType = strings.TrimSpace(sourceType)
sourceURI = strings.TrimSpace(sourceURI)
content = strings.TrimSpace(content)
if title == "" || len(title) > 256 {
return KnowledgeDocument{}, errors.New("文档标题格式无效")
}
if sourceType == "" {
sourceType = "text"
}
if sourceType != "text" && sourceType != "url" && sourceType != "import" {
return KnowledgeDocument{}, errors.New("source_type 无效")
}
if content == "" || len([]byte(content)) > 2<<20 {
return KnowledgeDocument{}, errors.New("文档正文不能为空且最多 2 MiB")
}
kb, err := s.GetKnowledgeBase(ctx, kbID)
if err != nil {
return KnowledgeDocument{}, err
}
chunks := ChunkText(content, kb.ChunkSize, kb.ChunkOverlap)
if len(chunks) == 0 {
return KnowledgeDocument{}, errors.New("文档没有可入库内容")
}
digest := sha256.Sum256([]byte(content))
hash := hex.EncodeToString(digest[:])
docID, err := newUUID()
if err != nil {
return KnowledgeDocument{}, err
}
tx, err := s.pool.Begin(ctx)
if err != nil {
return KnowledgeDocument{}, err
}
defer rollback(ctx, tx)
var doc KnowledgeDocument
err = tx.QueryRow(ctx, `INSERT INTO gateway.knowledge_documents(id,knowledge_base_id,title,source_type,source_uri,content,content_sha256,chunk_count,created_by) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9) RETURNING id::text,knowledge_base_id::text,title,source_type,source_uri,content_sha256,status,status_message,char_length(content),chunk_count,created_at,updated_at`, docID, kbID, title, sourceType, sourceURI, content, hash, len(chunks), actorID).Scan(&doc.ID, &doc.KnowledgeBaseID, &doc.Title, &doc.SourceType, &doc.SourceURI, &doc.ContentSHA256, &doc.Status, &doc.StatusMessage, &doc.CharCount, &doc.ChunkCount, &doc.CreatedAt, &doc.UpdatedAt)
if err != nil {
return KnowledgeDocument{}, err
}
embed := s.computeEmbeddings(ctx, kb, chunks)
chunkRows := make([][]any, 0, len(chunks))
chunkIDs := make([]string, 0, len(chunks))
embeddings := make([]string, 0, len(chunks))
for index, chunk := range chunks {
chunkID, idErr := newUUID()
if idErr != nil {
return KnowledgeDocument{}, idErr
}
chunkRows = append(chunkRows, []any{chunkID, kbID, docID, index, chunk})
// embed.embeddings 仅在成功算出向量时非 nil;降级(失败)时跳过回填。
if embed.embeddings != nil {
chunkIDs = append(chunkIDs, chunkID)
embeddings = append(embeddings, formatVector(embed.embeddings[index]))
}
}
// Bulk-copy all chunks in one statement instead of one INSERT per chunk;
// a 2 MiB document can split into thousands of chunks.
if _, err = tx.CopyFrom(ctx, pgx.Identifier{"gateway", "knowledge_chunks"}, []string{"id", "knowledge_base_id", "document_id", "chunk_index", "content"}, pgx.CopyFromRows(chunkRows)); err != nil {
return KnowledgeDocument{}, err
}
// pgx CopyFrom 对未知 OID(vector)列走二进制编码,直接随 COPY 传字面量会报
// "vector cannot have more than 16000 dimensions";因此向量在 COPY 之后用
// unnest 批量回填(同一事务内,单条 UPDATE)。
if embed.embeddings != nil {
// 降级时 embeddings 为空,跳过向量回填;否则空数组 UPDATE 也是无害空操作。
if _, err = tx.Exec(ctx, `UPDATE gateway.knowledge_chunks c SET embedding = v.vec::vector
FROM unnest($1::uuid[], $2::text[]) AS v(id, vec) WHERE c.id = v.id`, chunkIDs, embeddings); err != nil {
return KnowledgeDocument{}, err
}
}
if err = emit(ctx, tx, "knowledge_document.ready", "knowledge_document", docID, actorID, map[string]any{"knowledge_base_id": kbID, "chunk_count": len(chunks)}); err != nil {
return KnowledgeDocument{}, err
}
if embed.degraded {
if err = emit(ctx, tx, "knowledge_document.embedding_failed", "knowledge_document", docID, actorID, map[string]any{"knowledge_base_id": kbID, "chunk_count": len(chunks)}); err != nil {
return KnowledgeDocument{}, err
}
}
if err = tx.Commit(ctx); err != nil {
return KnowledgeDocument{}, err
}
return doc, nil
}
func (s *Service) ListKnowledgeDocuments(ctx context.Context, kbID string) ([]KnowledgeDocument, error) {
rows, err := s.pool.Query(ctx, `SELECT id::text,knowledge_base_id::text,title,source_type,source_uri,content_sha256,status,status_message,char_length(content),chunk_count,created_at,updated_at FROM gateway.knowledge_documents WHERE knowledge_base_id=$1 ORDER BY created_at DESC`, kbID)
if err != nil {
return nil, err
}
defer rows.Close()
items := []KnowledgeDocument{}
for rows.Next() {
var d KnowledgeDocument
if err = rows.Scan(&d.ID, &d.KnowledgeBaseID, &d.Title, &d.SourceType, &d.SourceURI, &d.ContentSHA256, &d.Status, &d.StatusMessage, &d.CharCount, &d.ChunkCount, &d.CreatedAt, &d.UpdatedAt); err != nil {
return nil, err
}
items = append(items, d)
}
return items, rows.Err()
}
func (s *Service) DeleteKnowledgeDocument(ctx context.Context, kbID, id, actorID string) error {
tx, err := s.pool.Begin(ctx)
if err != nil {
return err
}
defer rollback(ctx, tx)
tag, err := tx.Exec(ctx, `DELETE FROM gateway.knowledge_documents WHERE id=$1 AND knowledge_base_id=$2`, id, kbID)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrNotFound
}
if err = emit(ctx, tx, "knowledge_document.deleted", "knowledge_document", id, actorID, map[string]any{"knowledge_base_id": kbID}); err != nil {
return err
}
return tx.Commit(ctx)
}
func (s *Service) ReprocessKnowledgeDocument(ctx context.Context, kbID, id, actorID string) (KnowledgeDocument, error) {
kb, err := s.GetKnowledgeBase(ctx, kbID)
if err != nil {
return KnowledgeDocument{}, err
}
var content string
var doc KnowledgeDocument
err = s.pool.QueryRow(ctx, `SELECT id::text,knowledge_base_id::text,title,source_type,source_uri,content,content_sha256,status,status_message,char_length(content),chunk_count,created_at,updated_at FROM gateway.knowledge_documents WHERE id=$1 AND knowledge_base_id=$2`, id, kbID).Scan(&doc.ID, &doc.KnowledgeBaseID, &doc.Title, &doc.SourceType, &doc.SourceURI, &content, &doc.ContentSHA256, &doc.Status, &doc.StatusMessage, &doc.CharCount, &doc.ChunkCount, &doc.CreatedAt, &doc.UpdatedAt)
if errors.Is(err, pgx.ErrNoRows) {
return KnowledgeDocument{}, ErrNotFound
}
if err != nil {
return KnowledgeDocument{}, err
}
chunks := ChunkText(content, kb.ChunkSize, kb.ChunkOverlap)
if len(chunks) == 0 {
return KnowledgeDocument{}, errors.New("文档没有可入库内容")
}
tx, err := s.pool.Begin(ctx)
if err != nil {
return KnowledgeDocument{}, err
}
defer rollback(ctx, tx)
if _, err = tx.Exec(ctx, `DELETE FROM gateway.knowledge_chunks WHERE document_id=$1`, id); err != nil {
return KnowledgeDocument{}, err
}
embed := s.computeEmbeddings(ctx, kb, chunks)
for index, chunk := range chunks {
chunkID, idErr := newUUID()
if idErr != nil {
return KnowledgeDocument{}, idErr
}
if embed.embeddings != nil {
if _, err = tx.Exec(ctx, `INSERT INTO gateway.knowledge_chunks(id,knowledge_base_id,document_id,chunk_index,content,embedding) VALUES($1,$2,$3,$4,$5,$6::vector)`, chunkID, kbID, id, index, chunk, formatVector(embed.embeddings[index])); err != nil {
return KnowledgeDocument{}, err
}
} else {
if _, err = tx.Exec(ctx, `INSERT INTO gateway.knowledge_chunks(id,knowledge_base_id,document_id,chunk_index,content) VALUES($1,$2,$3,$4,$5)`, chunkID, kbID, id, index, chunk); err != nil {
return KnowledgeDocument{}, err
}
}
}
err = tx.QueryRow(ctx, `UPDATE gateway.knowledge_documents SET status='ready',status_message='',chunk_count=$3,updated_at=clock_timestamp() WHERE id=$1 AND knowledge_base_id=$2 RETURNING updated_at`, id, kbID, len(chunks)).Scan(&doc.UpdatedAt)
if err != nil {
return KnowledgeDocument{}, err
}
doc.Status, doc.StatusMessage, doc.ChunkCount = "ready", "", len(chunks)
if err = emit(ctx, tx, "knowledge_document.reprocessed", "knowledge_document", id, actorID, map[string]any{"knowledge_base_id": kbID, "chunk_count": len(chunks)}); err != nil {
return KnowledgeDocument{}, err
}
if embed.degraded {
if err = emit(ctx, tx, "knowledge_document.embedding_failed", "knowledge_document", id, actorID, map[string]any{"knowledge_base_id": kbID, "chunk_count": len(chunks)}); err != nil {
return KnowledgeDocument{}, err
}
}
if err = tx.Commit(ctx); err != nil {
return KnowledgeDocument{}, err
}
return doc, nil
}
func RuneCount(value string) int { return utf8.RuneCountInString(value) }
var _ Retriever = (*PostgreSQLRetriever)(nil)