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)