Files
ai-gateway-go/internal/workbench/knowledge.go
T
superidou 5759c1862e AI Gateway Go 0.10.0 源码快照 + 旗舰版需求规划报告
M0-M7 已完成:核心网关(身份/RBAC/TOTP/OIDC/SAML/Provider/配额/路由/内容策略/审计/定价)+ 资源市场(MCP/Skills/数字员工)。
含 22 个 PostgreSQL 迁移、管理端/门户端前端源码、OpenAPI 契约、部署 compose。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-08-12 11:45:54 +08:00

449 lines
15 KiB
Go

package workbench
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"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 NewPostgreSQLRetriever(service *Service) *PostgreSQLRetriever {
return &PostgreSQLRetriever{pool: service.pool}
}
func (r *PostgreSQLRetriever) Search(ctx context.Context, knowledgeBaseID, query string, topK int) ([]SearchHit, error) {
query = strings.TrimSpace(query)
if query == "" {
return nil, errors.New("检索词不能为空")
}
if len(query) > 6000 {
return nil, errors.New("检索词过长")
}
if topK < 1 {
topK = 4
}
if topK > 20 {
topK = 20
}
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"
}
if k.RetrievalMode != "postgres_fts" {
return errors.New("基线仅支持 postgres_fts 检索器")
}
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,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.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
}
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
}
chunkRows := make([][]any, 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})
}
// 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
}
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 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
}
for index, chunk := range chunks {
chunkID, idErr := newUUID()
if idErr != nil {
return KnowledgeDocument{}, idErr
}
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 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)