Files
ai-agent/internal/ai/rag/index.go
T
t 18c9354095 refactor: 将客服后端重构为宿主可嵌入模块
- 注入数据库、运行时配置、统一响应、文件存储和平台 AI 能力,补充业务读写工具与客户快捷操作契约。

- 移除模块内重复的组织、客户、工单、标签、技能、旧工作流、MCP 和迁移实现,将身份权限与业务主体交由宿主管理。

- 使用 libSQL 重构向量存储,并完善图片消息、访客身份、排队调度、企业微信和支持聊天页面。

- 统一 HTTP、DTO 与 WebSocket 的 snake_case 协议,补齐模块初始化、业务动作和公共载荷等回归测试。
2026-08-28 22:23:13 +08:00

327 lines
9.9 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package rag
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"log/slog"
"time"
"code.tczkiot.com/wlw/ai-agent/internal/ai"
ragchunk "code.tczkiot.com/wlw/ai-agent/internal/ai/rag/chunk"
"code.tczkiot.com/wlw/ai-agent/internal/ai/rag/vectordb"
"code.tczkiot.com/wlw/ai-agent/internal/models"
"code.tczkiot.com/wlw/ai-agent/internal/pkg/enums"
"code.tczkiot.com/wlw/ai-agent/internal/repositories"
"github.com/google/uuid"
"github.com/mlogclub/simple/sqls"
)
type ChunkingConfig struct {
Provider string
TargetTokens int
MaxTokens int
OverlapTokens int
EnableFallback bool
}
type index struct {
chunkConfig ChunkingConfig
registry *ragchunk.Registry
}
const knowledgeCollectionName = "knowledge_chunks"
var Index = &index{
chunkConfig: ChunkingConfig{
Provider: string(enums.KnowledgeChunkProviderStructured),
TargetTokens: 300,
MaxTokens: 400,
OverlapTokens: 40,
EnableFallback: true,
},
registry: ragchunk.NewDefaultRegistry(),
}
func (s *index) IndexDocumentByID(ctx context.Context, documentID int64) error {
document, err := s.loadDocumentByID(documentID)
if err != nil {
return err
}
return s.IndexDocument(ctx, *document)
}
func (s *index) IndexDocument(ctx context.Context, document models.KnowledgeDocument) error {
start := time.Now()
if err := s.markDocumentIndexPending(document.ID); err != nil {
slog.Error("Failed to mark knowledge document index as pending", "document_id", document.ID, "error", err)
}
fail := func(err error) error {
if updateErr := s.markDocumentIndexFailed(document.ID, err); updateErr != nil {
slog.Error("Failed to mark knowledge document index as failed", "document_id", document.ID, "error", updateErr)
}
return err
}
// TODO 这里每次都查询下知识库不太友好
knowledgeBase, err := s.loadDocumentKnowledgeBase(document)
if err != nil {
return fail(err)
}
vectors, chunkCount, err := s.runDocumentIndex(ctx, document, *knowledgeBase)
if err != nil {
return fail(err)
}
if err := s.markDocumentIndexIndexed(document.ID); err != nil {
slog.Error("Failed to mark knowledge document index as indexed", "document_id", document.ID, "error", err)
}
slog.Info("Document indexed successfully",
slog.Any("document_id", document.ID),
slog.Any("chunks_count", chunkCount),
slog.Any("vectors_count", len(vectors)),
slog.Any("time_taken", time.Since(start).String()),
)
return nil
}
func (s *index) IndexFAQByID(ctx context.Context, faqID int64) error {
faq, err := s.loadFAQByID(faqID)
if err != nil {
return err
}
if err := s.markFAQIndexPending(faq.ID); err != nil {
slog.Error("Failed to mark knowledge faq index as pending", "faq_id", faq.ID, "error", err)
}
fail := func(err error) error {
if updateErr := s.markFAQIndexFailed(faq.ID, err); updateErr != nil {
slog.Error("Failed to mark knowledge faq index as failed", "faq_id", faq.ID, "error", updateErr)
}
return err
}
knowledgeBase, err := s.loadFAQKnowledgeBase(*faq)
if err != nil {
return fail(err)
}
if err := s.runFAQIndex(ctx, *faq, *knowledgeBase); err != nil {
return fail(err)
}
if err := s.markFAQIndexIndexed(faq.ID); err != nil {
slog.Error("Failed to mark knowledge faq index as indexed", "faq_id", faq.ID, "error", err)
}
return nil
}
func (s *index) RemoveDocumentIndex(ctx context.Context, documentID int64) error {
chunks := repositories.KnowledgeChunkRepository.FindByDocumentID(sqls.DB(), documentID)
if len(chunks) == 0 {
return nil
}
if err := s.deleteChunkVectors(ctx, s.collectChunkVectorIDs(chunks)); err != nil {
slog.Error("Failed to delete vectors", "error", err)
}
if err := repositories.KnowledgeChunkRepository.DeleteByDocumentID(sqls.DB(), documentID); err != nil {
return fmt.Errorf("failed to delete chunks: %w", err)
}
slog.Info("Document index removed", "document_id", documentID, "chunks_removed", len(chunks))
return nil
}
func (s *index) RemoveFAQIndex(ctx context.Context, faqID int64) error {
chunks := repositories.KnowledgeChunkRepository.FindByFaqID(sqls.DB(), faqID)
if len(chunks) == 0 {
return nil
}
if err := s.deleteChunkVectors(ctx, s.collectChunkVectorIDs(chunks)); err != nil {
slog.Error("Failed to delete faq vectors", "error", err)
}
if err := repositories.KnowledgeChunkRepository.DeleteByFaqID(sqls.DB(), faqID); err != nil {
return fmt.Errorf("failed to delete faq chunks: %w", err)
}
slog.Info("FAQ index removed", "faq_id", faqID, "chunks_removed", len(chunks))
return nil
}
func (s *index) RemoveKnowledgeBaseIndex(ctx context.Context, knowledgeBaseID int64) error {
chunks := repositories.KnowledgeChunkRepository.FindByKnowledgeBaseID(sqls.DB(), knowledgeBaseID)
if len(chunks) == 0 {
return nil
}
if err := s.deleteChunkVectors(ctx, s.collectChunkVectorIDs(chunks)); err != nil {
slog.Error("Failed to delete knowledge base vectors", "knowledge_base_id", knowledgeBaseID, "error", err)
}
if err := repositories.KnowledgeChunkRepository.DeleteByKnowledgeBaseID(sqls.DB(), knowledgeBaseID); err != nil {
return fmt.Errorf("failed to delete chunks for knowledge base %d: %w", knowledgeBaseID, err)
}
slog.Info("Knowledge base index removed", "knowledge_base_id", knowledgeBaseID, "chunks_removed", len(chunks))
return nil
}
func (s *index) getCollectionName() string {
return knowledgeCollectionName
}
func buildKnowledgeChunkVectorID(knowledgeBaseID int64, documentID int64, chunkNo int) string {
raw := fmt.Sprintf("kb:%d:doc:%d:chunk:%d", knowledgeBaseID, documentID, chunkNo)
return uuid.NewSHA1(uuid.NameSpaceOID, []byte(raw)).String()
}
func buildKnowledgeFAQChunkVectorID(knowledgeBaseID int64, faqID int64, chunkNo int) string {
raw := fmt.Sprintf("kb:%d:faq:%d:chunk:%d", knowledgeBaseID, faqID, chunkNo)
return uuid.NewSHA1(uuid.NameSpaceOID, []byte(raw)).String()
}
func buildChunkContentHash(content string) string {
sum := sha256.Sum256([]byte(content))
return hex.EncodeToString(sum[:])
}
func firstPositiveInt(values ...int) int {
for _, value := range values {
if value > 0 {
return value
}
}
return 0
}
func firstNonEmptyString(values ...string) string {
for _, value := range values {
if value != "" {
return value
}
}
return ""
}
func (s *index) EnsureCollection(ctx context.Context) error {
dimension, err := ai.Embedding.GetDimension(ctx)
if err != nil {
return fmt.Errorf("failed to get embedding dimension: %w", err)
}
collectionName := s.getCollectionName()
provider := vectordb.GetProvider()
if provider == nil {
return fmt.Errorf("vectordb provider not initialized")
}
return s.ensureCollection(ctx, provider, collectionName, dimension)
}
func (s *index) RebuildKnowledgeBaseIndex(ctx context.Context, knowledgeBaseID int64) error {
knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), knowledgeBaseID)
if knowledgeBase == nil {
return fmt.Errorf("knowledge base not found: %d", knowledgeBaseID)
}
if err := s.resetKnowledgeBaseIndexStorage(ctx, knowledgeBaseID); err != nil {
return err
}
successCount := 0
failedCount := 0
if knowledgeBase.KnowledgeType == string(enums.KnowledgeBaseTypeFAQ) {
faqs := repositories.KnowledgeFAQRepository.Find(sqls.DB(), sqls.NewCnd().
Eq("knowledge_base_id", knowledgeBaseID).
Where("status != ?", enums.StatusDeleted))
if len(faqs) == 0 {
slog.Info("No faqs found in knowledge base, nothing to rebuild", "knowledge_base_id", knowledgeBaseID)
return nil
}
slog.Info("Rebuilding faq knowledge base index", "knowledge_base_id", knowledgeBaseID, "faq_count", len(faqs))
for _, faq := range faqs {
if err := s.IndexFAQByID(ctx, faq.ID); err != nil {
slog.Error("Failed to index faq", "faq_id", faq.ID, "error", err)
failedCount++
} else {
successCount++
}
}
} else {
documents := repositories.KnowledgeDocumentRepository.Find(sqls.DB(), sqls.NewCnd().
Eq("knowledge_base_id", knowledgeBaseID).
Where("status != ?", enums.StatusDeleted))
if len(documents) == 0 {
slog.Info("No documents found in knowledge base, nothing to rebuild", "knowledge_base_id", knowledgeBaseID)
return nil
}
documentIDs := make([]int64, 0, len(documents))
for _, doc := range documents {
documentIDs = append(documentIDs, doc.ID)
}
if err := s.markKnowledgeBaseDocumentsIndexPending(knowledgeBaseID, documentIDs); err != nil {
slog.Error("Failed to mark knowledge base documents index as pending", "knowledge_base_id", knowledgeBaseID, "error", err)
}
slog.Info("Rebuilding knowledge base index", "knowledge_base_id", knowledgeBaseID, "document_count", len(documents))
for _, doc := range documents {
if err := s.IndexDocumentByID(ctx, doc.ID); err != nil {
slog.Error("Failed to index document", "document_id", doc.ID, "error", err)
failedCount++
} else {
successCount++
}
}
}
slog.Info("Knowledge base index rebuild completed",
"knowledge_base_id", knowledgeBaseID,
"success_count", successCount,
"failed_count", failedCount)
return nil
}
func buildFAQChunkContent(faq models.KnowledgeFAQ) string {
parts := []string{fmt.Sprintf("问题:%s", faq.Question)}
var similarQuestions []string
if faq.SimilarQuestions != "" {
_ = json.Unmarshal([]byte(faq.SimilarQuestions), &similarQuestions)
}
if len(similarQuestions) > 0 {
parts = append(parts, fmt.Sprintf("相似问:%s", joinSimilarQuestions(similarQuestions)))
}
parts = append(parts, fmt.Sprintf("回答:%s", faq.Answer))
content := ""
for _, part := range parts {
if part == "" {
continue
}
if content != "" {
content += "\n"
}
content += part
}
return content
}
func joinSimilarQuestions(items []string) string {
result := ""
for _, item := range items {
if item == "" {
continue
}
if result != "" {
result += ""
}
result += item
}
return result
}
func (s *index) resetKnowledgeBaseIndexStorage(ctx context.Context, knowledgeBaseID int64) error {
chunks := repositories.KnowledgeChunkRepository.FindByKnowledgeBaseID(sqls.DB(), knowledgeBaseID)
return s.cleanupKnowledgeBaseChunks(ctx, knowledgeBaseID, chunks)
}