Init
This commit is contained in:
@@ -0,0 +1,432 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto"
|
||||
"cs-agent/internal/pkg/dto/request"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/pkg/errorsx"
|
||||
"cs-agent/internal/repositories"
|
||||
)
|
||||
|
||||
type answer struct {
|
||||
}
|
||||
|
||||
var Answer = &answer{}
|
||||
|
||||
func (s *answer) DebugSearch(ctx context.Context, req request.KnowledgeSearchRequest) (*response.KnowledgeSearchResponse, error) {
|
||||
if strings.TrimSpace(req.Question) == "" {
|
||||
return nil, errorsx.InvalidParam("问题不能为空")
|
||||
}
|
||||
startedAt := time.Now()
|
||||
results, err := s.retrieve(req, ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
respResults := make([]response.KnowledgeSearchResult, 0, len(results))
|
||||
for _, item := range results {
|
||||
respResults = append(respResults, response.KnowledgeSearchResult{
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
ChunkID: item.ChunkID,
|
||||
DocumentID: item.DocumentID,
|
||||
DocumentTitle: item.DocumentTitle,
|
||||
FaqID: item.FaqID,
|
||||
FaqQuestion: item.FaqQuestion,
|
||||
ChunkNo: item.ChunkNo,
|
||||
Title: item.Title,
|
||||
SectionPath: item.SectionPath,
|
||||
Content: item.Content,
|
||||
Score: float64(item.Score),
|
||||
})
|
||||
}
|
||||
|
||||
return &response.KnowledgeSearchResponse{
|
||||
Question: req.Question,
|
||||
Results: respResults,
|
||||
HitCount: len(respResults),
|
||||
LatencyMs: time.Since(startedAt).Milliseconds(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *answer) DebugAnswer(ctx context.Context, req request.KnowledgeAnswerRequest, operator *dto.AuthPrincipal) (*response.KnowledgeAnswerResponse, error) {
|
||||
if strings.TrimSpace(req.Question) == "" {
|
||||
return nil, errorsx.InvalidParam("问题不能为空")
|
||||
}
|
||||
startedAt := time.Now()
|
||||
|
||||
retrieveStartedAt := time.Now()
|
||||
results, err := s.retrieve(request.KnowledgeSearchRequest{
|
||||
KnowledgeBaseIDs: req.KnowledgeBaseIDs,
|
||||
Question: req.Question,
|
||||
TopK: req.TopK,
|
||||
ScoreThreshold: req.ScoreThreshold,
|
||||
RerankLimit: req.RerankLimit,
|
||||
}, ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
retrieveMs := time.Since(retrieveStartedAt).Milliseconds()
|
||||
knowledgeBase := s.resolveAnswerKnowledgeBase(req.KnowledgeBaseIDs, results)
|
||||
contextResults := buildContextHits(Retrieve.SelectContextResults(results, 4000))
|
||||
|
||||
hits := make([]response.KnowledgeSearchResult, 0, len(results))
|
||||
topScore := 0.0
|
||||
for i, item := range results {
|
||||
score := float64(item.Score)
|
||||
if i == 0 {
|
||||
topScore = score
|
||||
}
|
||||
hits = append(hits, response.KnowledgeSearchResult{
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
ChunkID: item.ChunkID,
|
||||
DocumentID: item.DocumentID,
|
||||
DocumentTitle: item.DocumentTitle,
|
||||
FaqID: item.FaqID,
|
||||
FaqQuestion: item.FaqQuestion,
|
||||
ChunkNo: item.ChunkNo,
|
||||
Title: item.Title,
|
||||
SectionPath: item.SectionPath,
|
||||
Content: item.Content,
|
||||
Score: score,
|
||||
})
|
||||
}
|
||||
citations := buildKnowledgeCitations(hits, 3)
|
||||
|
||||
answerMode := enums.KnowledgeAnswerMode(req.AnswerMode)
|
||||
if answerMode == 0 {
|
||||
if knowledgeBase != nil {
|
||||
answerMode = enums.KnowledgeAnswerMode(knowledgeBase.AnswerMode)
|
||||
}
|
||||
if answerMode == 0 {
|
||||
answerMode = enums.KnowledgeAnswerModeStrict
|
||||
}
|
||||
}
|
||||
|
||||
fallbackMode := enums.KnowledgeFallbackMode(req.FallbackMode)
|
||||
if fallbackMode == 0 {
|
||||
if knowledgeBase != nil {
|
||||
fallbackMode = enums.KnowledgeFallbackMode(knowledgeBase.FallbackMode)
|
||||
}
|
||||
if fallbackMode == 0 {
|
||||
fallbackMode = enums.KnowledgeFallbackModeNoAnswer
|
||||
}
|
||||
}
|
||||
|
||||
answerStatus := enums.KnowledgeAnswerStatusNormal
|
||||
answer := ""
|
||||
modelName := ""
|
||||
promptTokens := 0
|
||||
completionTokens := 0
|
||||
generateStartedAt := time.Now()
|
||||
|
||||
if len(hits) == 0 {
|
||||
answerStatus = enums.KnowledgeAnswerStatusNoAnswer
|
||||
answer = buildFallbackAnswer(fallbackMode)
|
||||
} else {
|
||||
contextText := Retrieve.BuildContext(ctx, results, 4000)
|
||||
systemPrompt := buildAnswerSystemPrompt(answerMode)
|
||||
userPrompt := fmt.Sprintf("用户问题:%s\n\n参考资料:\n%s", req.Question, contextText)
|
||||
llmResult, llmErr := ai.LLM.Chat(ctx, systemPrompt, userPrompt)
|
||||
if llmErr != nil {
|
||||
answerStatus = enums.KnowledgeAnswerStatusFallback
|
||||
answer = buildFallbackAnswer(fallbackMode)
|
||||
} else {
|
||||
answer = llmResult.Content
|
||||
modelName = llmResult.ModelName
|
||||
promptTokens = llmResult.PromptTokens
|
||||
completionTokens = llmResult.CompletionTokens
|
||||
if strings.TrimSpace(answer) == "" {
|
||||
answerStatus = enums.KnowledgeAnswerStatusFallback
|
||||
answer = buildFallbackAnswer(fallbackMode)
|
||||
}
|
||||
}
|
||||
}
|
||||
generateMs := time.Since(generateStartedAt).Milliseconds()
|
||||
rerankLimit := 0
|
||||
chunkProvider := ""
|
||||
chunkTargetTokens := 0
|
||||
chunkMaxTokens := 0
|
||||
chunkOverlapTokens := 0
|
||||
if knowledgeBase != nil {
|
||||
rerankLimit = resolveRerankLimit(req.RerankLimit, knowledgeBase.DefaultRerankLimit)
|
||||
chunkProvider = knowledgeBase.ChunkProvider
|
||||
chunkTargetTokens = knowledgeBase.ChunkTargetTokens
|
||||
chunkMaxTokens = knowledgeBase.ChunkMaxTokens
|
||||
chunkOverlapTokens = knowledgeBase.ChunkOverlapTokens
|
||||
}
|
||||
|
||||
logItem, err := RetrieveLog.CreateRetrieveLog(&CreateRetrieveLogRequest{
|
||||
KnowledgeBaseID: firstKnowledgeBaseID(req.KnowledgeBaseIDs),
|
||||
Channel: defaultRetrieveChannel(req.Channel),
|
||||
Scene: defaultRetrieveScene(req.Scene),
|
||||
SessionID: req.SessionID,
|
||||
ConversationID: req.ConversationID,
|
||||
Question: req.Question,
|
||||
RewriteQuestion: "",
|
||||
Answer: answer,
|
||||
AnswerStatus: int(answerStatus),
|
||||
ChunkProvider: chunkProvider,
|
||||
ChunkTargetTokens: chunkTargetTokens,
|
||||
ChunkMaxTokens: chunkMaxTokens,
|
||||
ChunkOverlapTokens: chunkOverlapTokens,
|
||||
RerankEnabled: rerankLimit > 0,
|
||||
RerankLimit: rerankLimit,
|
||||
Hits: hits,
|
||||
UsedHits: contextResults,
|
||||
Citations: citations,
|
||||
LatencyMs: time.Since(startedAt).Milliseconds(),
|
||||
RetrieveMs: retrieveMs,
|
||||
GenerateMs: generateMs,
|
||||
PromptTokens: promptTokens,
|
||||
CompletionTokens: completionTokens,
|
||||
ModelName: modelName,
|
||||
}, operator)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &response.KnowledgeAnswerResponse{
|
||||
Question: req.Question,
|
||||
Answer: answer,
|
||||
AnswerStatus: int(answerStatus),
|
||||
AnswerStatusName: getAnswerStatusName(answerStatus),
|
||||
Citations: citations,
|
||||
Hits: hits,
|
||||
HitCount: len(hits),
|
||||
TopScore: topScore,
|
||||
LatencyMs: time.Since(startedAt).Milliseconds(),
|
||||
RetrieveMs: retrieveMs,
|
||||
GenerateMs: generateMs,
|
||||
PromptTokens: promptTokens,
|
||||
CompletionTokens: completionTokens,
|
||||
ModelName: modelName,
|
||||
RetrieveLogID: logItem.ID,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildContextHits(results []RetrieveResult) []response.KnowledgeSearchResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
hits := make([]response.KnowledgeSearchResult, 0, len(results))
|
||||
for _, item := range results {
|
||||
hits = append(hits, response.KnowledgeSearchResult{
|
||||
KnowledgeBaseID: item.KnowledgeBaseID,
|
||||
ChunkID: item.ChunkID,
|
||||
DocumentID: item.DocumentID,
|
||||
DocumentTitle: item.DocumentTitle,
|
||||
FaqID: item.FaqID,
|
||||
FaqQuestion: item.FaqQuestion,
|
||||
ChunkNo: item.ChunkNo,
|
||||
Title: item.Title,
|
||||
SectionPath: item.SectionPath,
|
||||
Content: item.Content,
|
||||
Score: float64(item.Score),
|
||||
})
|
||||
}
|
||||
return hits
|
||||
}
|
||||
|
||||
func buildKnowledgeCitations(hits []response.KnowledgeSearchResult, limit int) []response.KnowledgeCitation {
|
||||
if len(hits) == 0 || limit <= 0 {
|
||||
return nil
|
||||
}
|
||||
citations := make([]response.KnowledgeCitation, 0, limit)
|
||||
seen := make(map[string]struct{})
|
||||
for _, item := range hits {
|
||||
key := fmt.Sprintf("%d|%d|%s|%d", item.DocumentID, item.FaqID, item.SectionPath, item.ChunkNo)
|
||||
if _, ok := seen[key]; ok {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
citations = append(citations, response.KnowledgeCitation{
|
||||
DocumentID: item.DocumentID,
|
||||
DocumentTitle: item.DocumentTitle,
|
||||
FaqID: item.FaqID,
|
||||
FaqQuestion: item.FaqQuestion,
|
||||
ChunkNo: item.ChunkNo,
|
||||
Title: item.Title,
|
||||
SectionPath: item.SectionPath,
|
||||
Snippet: truncateCitationSnippet(item.Content, 120),
|
||||
Score: item.Score,
|
||||
})
|
||||
if len(citations) >= limit {
|
||||
break
|
||||
}
|
||||
}
|
||||
return citations
|
||||
}
|
||||
|
||||
func truncateCitationSnippet(text string, limit int) string {
|
||||
if limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(strings.TrimSpace(text))
|
||||
if len(runes) <= limit {
|
||||
return string(runes)
|
||||
}
|
||||
return string(runes[:limit]) + "..."
|
||||
}
|
||||
|
||||
func (s *answer) BuildDocumentIndex(ctx context.Context, documentID int64) error {
|
||||
return Index.IndexDocumentByID(ctx, documentID)
|
||||
}
|
||||
|
||||
func (s *answer) retrieve(req request.KnowledgeSearchRequest, ctx context.Context) ([]RetrieveResult, error) {
|
||||
if len(normalizeKnowledgeBaseIDs(req.KnowledgeBaseIDs)) == 0 {
|
||||
return nil, errorsx.InvalidParam("知识库不能为空")
|
||||
}
|
||||
knowledgeBases := s.loadKnowledgeBases(req.KnowledgeBaseIDs)
|
||||
|
||||
results, err := Retrieve.Retrieve(ctx, RetrieveRequest{
|
||||
KnowledgeBaseIDs: req.KnowledgeBaseIDs,
|
||||
Query: req.Question,
|
||||
TopK: req.TopK,
|
||||
ScoreThreshold: req.ScoreThreshold,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
defaultRerankLimit := resolveDefaultRerankLimit(knowledgeBases)
|
||||
rerankLimit := resolveRerankLimit(req.RerankLimit, defaultRerankLimit)
|
||||
if rerankLimit > 0 && len(results) > rerankLimit {
|
||||
return Retrieve.RetrieveWithRerank(ctx, RetrieveRequest{
|
||||
KnowledgeBaseIDs: req.KnowledgeBaseIDs,
|
||||
Query: req.Question,
|
||||
TopK: req.TopK,
|
||||
ScoreThreshold: req.ScoreThreshold,
|
||||
}, rerankLimit)
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (s *answer) loadKnowledgeBases(knowledgeBaseIDs []int64) []models.KnowledgeBase {
|
||||
normalized := normalizeKnowledgeBaseIDs(knowledgeBaseIDs)
|
||||
if len(normalized) == 0 {
|
||||
return nil
|
||||
}
|
||||
items := repositories.KnowledgeBaseRepository.Find(sqls.DB(), sqls.NewCnd().In("id", normalized))
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
itemMap := make(map[int64]models.KnowledgeBase, len(items))
|
||||
for _, item := range items {
|
||||
itemMap[item.ID] = item
|
||||
}
|
||||
results := make([]models.KnowledgeBase, 0, len(normalized))
|
||||
for _, id := range normalized {
|
||||
if item, ok := itemMap[id]; ok {
|
||||
results = append(results, item)
|
||||
}
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func (s *answer) resolvePrimaryKnowledgeBase(knowledgeBaseIDs []int64) *models.KnowledgeBase {
|
||||
items := s.loadKnowledgeBases(knowledgeBaseIDs)
|
||||
for _, item := range items {
|
||||
return &item
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *answer) resolveAnswerKnowledgeBase(knowledgeBaseIDs []int64, results []RetrieveResult) *models.KnowledgeBase {
|
||||
items := s.loadKnowledgeBases(knowledgeBaseIDs)
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
if len(results) > 0 {
|
||||
for _, item := range items {
|
||||
if item.ID == results[0].KnowledgeBaseID {
|
||||
return &item
|
||||
}
|
||||
}
|
||||
}
|
||||
return &items[0]
|
||||
}
|
||||
|
||||
func firstKnowledgeBaseID(ids []int64) int64 {
|
||||
normalized := normalizeKnowledgeBaseIDs(ids)
|
||||
if len(normalized) == 0 {
|
||||
return 0
|
||||
}
|
||||
return normalized[0]
|
||||
}
|
||||
|
||||
func resolveRerankLimit(requestLimit int, defaultLimit int) int {
|
||||
if requestLimit > 0 {
|
||||
return requestLimit
|
||||
}
|
||||
if defaultLimit > 0 {
|
||||
return defaultLimit
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func resolveDefaultRerankLimit(items []models.KnowledgeBase) int {
|
||||
limit := 0
|
||||
for _, item := range items {
|
||||
if item.DefaultRerankLimit > limit {
|
||||
limit = item.DefaultRerankLimit
|
||||
}
|
||||
}
|
||||
return limit
|
||||
}
|
||||
|
||||
func buildAnswerSystemPrompt(answerMode enums.KnowledgeAnswerMode) string {
|
||||
if answerMode == enums.KnowledgeAnswerModeAssist {
|
||||
return "你是客服知识库助手。请优先依据提供的知识片段回答,可以做轻度归纳,但不要编造未提供的事实。"
|
||||
}
|
||||
return "你是严格的客服知识库助手。只能依据提供的知识片段回答;如果资料不足,请明确说明知识库暂无明确信息。"
|
||||
}
|
||||
|
||||
func buildFallbackAnswer(fallbackMode enums.KnowledgeFallbackMode) string {
|
||||
switch fallbackMode {
|
||||
case enums.KnowledgeFallbackModeSuggestRetry:
|
||||
return "当前知识库里没有找到足够明确的信息,你可以换个更具体的问法再试一次。"
|
||||
case enums.KnowledgeFallbackModeTransferHuman:
|
||||
return "当前知识库里没有找到足够明确的信息,建议转人工进一步处理。"
|
||||
default:
|
||||
return "当前知识库暂无明确信息。"
|
||||
}
|
||||
}
|
||||
|
||||
func defaultRetrieveChannel(channel string) string {
|
||||
if strings.TrimSpace(channel) == "" {
|
||||
return string(enums.KnowledgeRetrieveChannelDebug)
|
||||
}
|
||||
return channel
|
||||
}
|
||||
|
||||
func defaultRetrieveScene(scene string) string {
|
||||
if strings.TrimSpace(scene) == "" {
|
||||
return string(enums.KnowledgeRetrieveSceneQA)
|
||||
}
|
||||
return scene
|
||||
}
|
||||
|
||||
func getAnswerStatusName(status enums.KnowledgeAnswerStatus) string {
|
||||
switch status {
|
||||
case enums.KnowledgeAnswerStatusNormal:
|
||||
return "正常"
|
||||
case enums.KnowledgeAnswerStatusNoAnswer:
|
||||
return "无答案"
|
||||
case enums.KnowledgeAnswerStatusFallback:
|
||||
return "兜底"
|
||||
case enums.KnowledgeAnswerStatusBlocked:
|
||||
return "风控拦截"
|
||||
default:
|
||||
return "未知"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
func TestBuildFallbackAnswer(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
mode enums.KnowledgeFallbackMode
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "no answer",
|
||||
mode: enums.KnowledgeFallbackModeNoAnswer,
|
||||
expected: "当前知识库暂无明确信息。",
|
||||
},
|
||||
{
|
||||
name: "suggest retry",
|
||||
mode: enums.KnowledgeFallbackModeSuggestRetry,
|
||||
expected: "当前知识库里没有找到足够明确的信息,你可以换个更具体的问法再试一次。",
|
||||
},
|
||||
{
|
||||
name: "transfer human",
|
||||
mode: enums.KnowledgeFallbackModeTransferHuman,
|
||||
expected: "当前知识库里没有找到足够明确的信息,建议转人工进一步处理。",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := buildFallbackAnswer(tt.mode); got != tt.expected {
|
||||
t.Fatalf("%s: expected %q, got %q", tt.name, tt.expected, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetAnswerStatusName(t *testing.T) {
|
||||
if got := getAnswerStatusName(enums.KnowledgeAnswerStatusNoAnswer); got != "无答案" {
|
||||
t.Fatalf("expected no-answer label, got %q", got)
|
||||
}
|
||||
if got := getAnswerStatusName(enums.KnowledgeAnswerStatusFallback); got != "兜底" {
|
||||
t.Fatalf("expected fallback label, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveRerankLimit(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
requestLimit int
|
||||
defaultLimit int
|
||||
expected int
|
||||
}{
|
||||
{
|
||||
name: "request overrides default",
|
||||
requestLimit: 3,
|
||||
defaultLimit: 5,
|
||||
expected: 3,
|
||||
},
|
||||
{
|
||||
name: "default used when request missing",
|
||||
requestLimit: 0,
|
||||
defaultLimit: 5,
|
||||
expected: 5,
|
||||
},
|
||||
{
|
||||
name: "zero when both missing",
|
||||
requestLimit: 0,
|
||||
defaultLimit: 0,
|
||||
expected: 0,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := resolveRerankLimit(tt.requestLimit, tt.defaultLimit); got != tt.expected {
|
||||
t.Fatalf("%s: expected %d, got %d", tt.name, tt.expected, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveDefaultRerankLimit(t *testing.T) {
|
||||
items := []models.KnowledgeBase{
|
||||
{ID: 11, DefaultRerankLimit: 3},
|
||||
{ID: 22, DefaultRerankLimit: 7},
|
||||
{ID: 33, DefaultRerankLimit: 5},
|
||||
}
|
||||
|
||||
if got := resolveDefaultRerankLimit(items); got != 7 {
|
||||
t.Fatalf("expected max rerank limit 7, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildKnowledgeCitations(t *testing.T) {
|
||||
hits := []response.KnowledgeSearchResult{
|
||||
{
|
||||
DocumentID: 11,
|
||||
DocumentTitle: "退款手册",
|
||||
ChunkNo: 0,
|
||||
Title: "退款说明",
|
||||
SectionPath: "售后 > 退款说明",
|
||||
Content: "退款申请提交后,预计1-3个工作日到账。",
|
||||
Score: 0.91,
|
||||
},
|
||||
{
|
||||
DocumentID: 11,
|
||||
DocumentTitle: "退款手册",
|
||||
ChunkNo: 0,
|
||||
Title: "退款说明",
|
||||
SectionPath: "售后 > 退款说明",
|
||||
Content: "重复内容",
|
||||
Score: 0.89,
|
||||
},
|
||||
}
|
||||
|
||||
citations := buildKnowledgeCitations(hits, 3)
|
||||
if len(citations) != 1 {
|
||||
t.Fatalf("expected 1 citation, got %d", len(citations))
|
||||
}
|
||||
if citations[0].DocumentID != 11 {
|
||||
t.Fatalf("expected document id 11, got %d", citations[0].DocumentID)
|
||||
}
|
||||
if citations[0].SectionPath != "售后 > 退款说明" {
|
||||
t.Fatalf("unexpected section path: %q", citations[0].SectionPath)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"context"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
type fixedProvider struct{}
|
||||
|
||||
func NewFixedProvider() Provider {
|
||||
return &fixedProvider{}
|
||||
}
|
||||
|
||||
func (p *fixedProvider) Name() string {
|
||||
return string(enums.KnowledgeChunkProviderFixed)
|
||||
}
|
||||
|
||||
func (p *fixedProvider) Supports(contentType enums.KnowledgeDocumentContentType) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (p *fixedProvider) Chunk(ctx context.Context, req *ChunkRequest) ([]ChunkResult, error) {
|
||||
text := req.PlainText
|
||||
if text == "" {
|
||||
text = req.Content
|
||||
}
|
||||
parts := splitPlainText(text, req.Options)
|
||||
results := make([]ChunkResult, 0, len(parts))
|
||||
for i, part := range parts {
|
||||
results = append(results, ChunkResult{
|
||||
ChunkNo: i,
|
||||
Title: req.DocumentTitle,
|
||||
Content: part,
|
||||
ChunkType: enums.KnowledgeChunkTypeText,
|
||||
SectionPath: req.DocumentTitle,
|
||||
CharCount: len([]rune(part)),
|
||||
TokenCount: estimateTokenCount(part),
|
||||
Metadata: map[string]any{
|
||||
"provider": enums.KnowledgeChunkProviderFixed,
|
||||
},
|
||||
})
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"context"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
type Provider interface {
|
||||
Name() string
|
||||
Supports(contentType enums.KnowledgeDocumentContentType) bool
|
||||
Chunk(ctx context.Context, req *ChunkRequest) ([]ChunkResult, error)
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"context"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
type Registry struct {
|
||||
providers map[string]Provider
|
||||
}
|
||||
|
||||
func NewRegistry() *Registry {
|
||||
return &Registry{
|
||||
providers: make(map[string]Provider),
|
||||
}
|
||||
}
|
||||
|
||||
func NewDefaultRegistry() *Registry {
|
||||
r := NewRegistry()
|
||||
r.Register(NewFixedProvider())
|
||||
r.Register(NewStructuredProvider())
|
||||
return r
|
||||
}
|
||||
|
||||
func (r *Registry) Register(p Provider) {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
r.providers[p.Name()] = p
|
||||
}
|
||||
|
||||
func (r *Registry) Get(name string) Provider {
|
||||
if name == "" {
|
||||
return nil
|
||||
}
|
||||
return r.providers[name]
|
||||
}
|
||||
|
||||
func (r *Registry) Resolve(name string, contentType enums.KnowledgeDocumentContentType) Provider {
|
||||
if p := r.Get(name); p != nil && p.Supports(contentType) {
|
||||
return p
|
||||
}
|
||||
if p := r.Get(string(enums.KnowledgeChunkProviderStructured)); p != nil && p.Supports(contentType) {
|
||||
return p
|
||||
}
|
||||
return r.Get(string(enums.KnowledgeChunkProviderFixed))
|
||||
}
|
||||
|
||||
func (r *Registry) Chunk(ctx context.Context, req *ChunkRequest) ([]ChunkResult, error) {
|
||||
if req == nil {
|
||||
return nil, fmt.Errorf("chunk request is nil")
|
||||
}
|
||||
provider := r.Resolve(req.Options.Provider, req.ContentType)
|
||||
if provider == nil {
|
||||
return nil, fmt.Errorf("chunk provider not found")
|
||||
}
|
||||
return provider.Chunk(ctx, req)
|
||||
}
|
||||
@@ -0,0 +1,261 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"context"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"strings"
|
||||
|
||||
"github.com/gomarkdown/markdown"
|
||||
"golang.org/x/net/html"
|
||||
)
|
||||
|
||||
type structuredProvider struct{}
|
||||
|
||||
type contentBlock struct {
|
||||
Type string
|
||||
Level int
|
||||
Text string
|
||||
Title string
|
||||
SectionPath string
|
||||
}
|
||||
|
||||
func NewStructuredProvider() Provider {
|
||||
return &structuredProvider{}
|
||||
}
|
||||
|
||||
func (p *structuredProvider) Name() string {
|
||||
return string(enums.KnowledgeChunkProviderStructured)
|
||||
}
|
||||
|
||||
func (p *structuredProvider) Supports(contentType enums.KnowledgeDocumentContentType) bool {
|
||||
switch contentType {
|
||||
case enums.KnowledgeDocumentContentTypeHTML, enums.KnowledgeDocumentContentTypeMarkdown:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (p *structuredProvider) Chunk(ctx context.Context, req *ChunkRequest) ([]ChunkResult, error) {
|
||||
content := req.Content
|
||||
if req.ContentType == enums.KnowledgeDocumentContentTypeMarkdown {
|
||||
content = string(markdown.ToHTML([]byte(content), nil, nil))
|
||||
}
|
||||
|
||||
blocks := parseStructuredBlocks(content, req.DocumentTitle)
|
||||
if len(blocks) == 0 {
|
||||
return NewFixedProvider().Chunk(ctx, req)
|
||||
}
|
||||
|
||||
results := make([]ChunkResult, 0)
|
||||
chunkNo := 0
|
||||
for _, block := range blocks {
|
||||
parts := splitPlainText(block.Text, req.Options)
|
||||
for _, part := range parts {
|
||||
if part == "" {
|
||||
continue
|
||||
}
|
||||
results = append(results, ChunkResult{
|
||||
ChunkNo: chunkNo,
|
||||
Title: block.Title,
|
||||
Content: part,
|
||||
ChunkType: mapBlockType(block.Type),
|
||||
SectionPath: block.SectionPath,
|
||||
CharCount: len([]rune(part)),
|
||||
TokenCount: estimateTokenCount(part),
|
||||
Metadata: map[string]any{
|
||||
"provider": enums.KnowledgeChunkProviderStructured,
|
||||
"blockType": block.Type,
|
||||
"sectionPath": block.SectionPath,
|
||||
"sectionTitle": block.Title,
|
||||
},
|
||||
})
|
||||
chunkNo++
|
||||
}
|
||||
}
|
||||
if len(results) == 0 {
|
||||
return NewFixedProvider().Chunk(ctx, req)
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func parseStructuredBlocks(content string, documentTitle string) []contentBlock {
|
||||
content = strings.TrimSpace(content)
|
||||
if content == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
parent := &html.Node{Type: html.ElementNode, Data: "div"}
|
||||
nodes, err := html.ParseFragment(strings.NewReader(content), parent)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
var blocks []contentBlock
|
||||
headings := make([]string, 0)
|
||||
var walk func(node *html.Node)
|
||||
walk = func(node *html.Node) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
if node.Type == html.ElementNode {
|
||||
switch node.Data {
|
||||
case "h1", "h2", "h3", "h4", "h5", "h6":
|
||||
title := normalizeText(nodeText(node))
|
||||
if title != "" {
|
||||
level := int(node.Data[1] - '0')
|
||||
if level <= 0 {
|
||||
level = 1
|
||||
}
|
||||
headings = updateHeadingPath(headings, level, title)
|
||||
}
|
||||
return
|
||||
case "p":
|
||||
appendBlock(&blocks, "paragraph", normalizeText(nodeText(node)), currentTitle(headings, documentTitle), strings.Join(headings, " > "))
|
||||
return
|
||||
case "ul", "ol":
|
||||
appendBlock(&blocks, "list", normalizeText(listText(node)), currentTitle(headings, documentTitle), strings.Join(headings, " > "))
|
||||
return
|
||||
case "table":
|
||||
appendBlock(&blocks, "table", normalizeText(tableText(node)), currentTitle(headings, documentTitle), strings.Join(headings, " > "))
|
||||
return
|
||||
case "pre", "code":
|
||||
appendBlock(&blocks, "code", normalizeText(nodeText(node)), currentTitle(headings, documentTitle), strings.Join(headings, " > "))
|
||||
return
|
||||
}
|
||||
}
|
||||
for child := node.FirstChild; child != nil; child = child.NextSibling {
|
||||
walk(child)
|
||||
}
|
||||
}
|
||||
for _, node := range nodes {
|
||||
walk(node)
|
||||
}
|
||||
return blocks
|
||||
}
|
||||
|
||||
func appendBlock(blocks *[]contentBlock, blockType string, text string, title string, sectionPath string) {
|
||||
text = normalizeText(text)
|
||||
if text == "" {
|
||||
return
|
||||
}
|
||||
if sectionPath == "" {
|
||||
sectionPath = title
|
||||
}
|
||||
*blocks = append(*blocks, contentBlock{
|
||||
Type: blockType,
|
||||
Text: text,
|
||||
Title: title,
|
||||
SectionPath: sectionPath,
|
||||
})
|
||||
}
|
||||
|
||||
func updateHeadingPath(headings []string, level int, title string) []string {
|
||||
if level <= 0 {
|
||||
level = 1
|
||||
}
|
||||
if len(headings) >= level {
|
||||
headings = headings[:level-1]
|
||||
}
|
||||
headings = append(headings, title)
|
||||
return headings
|
||||
}
|
||||
|
||||
func currentTitle(headings []string, documentTitle string) string {
|
||||
if len(headings) == 0 {
|
||||
return documentTitle
|
||||
}
|
||||
return headings[len(headings)-1]
|
||||
}
|
||||
|
||||
func mapBlockType(blockType string) enums.KnowledgeChunkType {
|
||||
switch blockType {
|
||||
case "table":
|
||||
return enums.KnowledgeChunkTypeTable
|
||||
case "code":
|
||||
return enums.KnowledgeChunkTypeCode
|
||||
default:
|
||||
return enums.KnowledgeChunkTypeText
|
||||
}
|
||||
}
|
||||
|
||||
func nodeText(node *html.Node) string {
|
||||
if node == nil {
|
||||
return ""
|
||||
}
|
||||
var builder strings.Builder
|
||||
writeNodeText(&builder, node)
|
||||
return builder.String()
|
||||
}
|
||||
|
||||
func writeNodeText(builder *strings.Builder, node *html.Node) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
switch node.Type {
|
||||
case html.TextNode:
|
||||
builder.WriteString(node.Data)
|
||||
case html.ElementNode:
|
||||
if shouldSeparate(node.Data) {
|
||||
builder.WriteByte(' ')
|
||||
}
|
||||
}
|
||||
for child := node.FirstChild; child != nil; child = child.NextSibling {
|
||||
writeNodeText(builder, child)
|
||||
}
|
||||
if node.Type == html.ElementNode && shouldSeparate(node.Data) {
|
||||
builder.WriteByte(' ')
|
||||
}
|
||||
}
|
||||
|
||||
func shouldSeparate(tag string) bool {
|
||||
switch tag {
|
||||
case "p", "div", "br", "li", "ul", "ol", "blockquote", "pre", "table", "tr", "td", "th", "h1", "h2", "h3", "h4", "h5", "h6":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func listText(node *html.Node) string {
|
||||
items := make([]string, 0)
|
||||
for child := node.FirstChild; child != nil; child = child.NextSibling {
|
||||
if child.Type == html.ElementNode && child.Data == "li" {
|
||||
item := normalizeText(nodeText(child))
|
||||
if item != "" {
|
||||
items = append(items, item)
|
||||
}
|
||||
}
|
||||
}
|
||||
return strings.Join(items, " ")
|
||||
}
|
||||
|
||||
func tableText(node *html.Node) string {
|
||||
rows := make([]string, 0)
|
||||
var walk func(*html.Node)
|
||||
walk = func(n *html.Node) {
|
||||
if n == nil {
|
||||
return
|
||||
}
|
||||
if n.Type == html.ElementNode && n.Data == "tr" {
|
||||
cells := make([]string, 0)
|
||||
for child := n.FirstChild; child != nil; child = child.NextSibling {
|
||||
if child.Type == html.ElementNode && (child.Data == "td" || child.Data == "th") {
|
||||
cell := normalizeText(nodeText(child))
|
||||
if cell != "" {
|
||||
cells = append(cells, cell)
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(cells) > 0 {
|
||||
rows = append(rows, strings.Join(cells, " | "))
|
||||
}
|
||||
return
|
||||
}
|
||||
for child := n.FirstChild; child != nil; child = child.NextSibling {
|
||||
walk(child)
|
||||
}
|
||||
}
|
||||
walk(node)
|
||||
return strings.Join(rows, " ")
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package chunk
|
||||
|
||||
import "cs-agent/internal/pkg/enums"
|
||||
|
||||
type ChunkRequest struct {
|
||||
KnowledgeBaseID int64
|
||||
DocumentID int64
|
||||
DocumentTitle string
|
||||
ContentType enums.KnowledgeDocumentContentType
|
||||
Content string
|
||||
PlainText string
|
||||
Options ChunkOptions
|
||||
}
|
||||
|
||||
type ChunkOptions struct {
|
||||
Provider string
|
||||
TargetTokens int
|
||||
MaxTokens int
|
||||
OverlapTokens int
|
||||
EnableFallback bool
|
||||
}
|
||||
|
||||
type ChunkResult struct {
|
||||
ChunkNo int
|
||||
Title string
|
||||
Content string
|
||||
ChunkType enums.KnowledgeChunkType
|
||||
SectionPath string
|
||||
CharCount int
|
||||
TokenCount int
|
||||
Metadata map[string]any
|
||||
}
|
||||
@@ -0,0 +1,218 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"encoding/hex"
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTargetTokens = 300
|
||||
defaultMaxTokens = 400
|
||||
defaultOverlapTokens = 40
|
||||
)
|
||||
|
||||
func normalizeOptions(opts ChunkOptions) ChunkOptions {
|
||||
if opts.TargetTokens <= 0 {
|
||||
opts.TargetTokens = defaultTargetTokens
|
||||
}
|
||||
if opts.MaxTokens <= 0 {
|
||||
opts.MaxTokens = defaultMaxTokens
|
||||
}
|
||||
if opts.MaxTokens < opts.TargetTokens {
|
||||
opts.MaxTokens = opts.TargetTokens
|
||||
}
|
||||
if opts.OverlapTokens < 0 {
|
||||
opts.OverlapTokens = 0
|
||||
}
|
||||
if opts.OverlapTokens == 0 {
|
||||
opts.OverlapTokens = defaultOverlapTokens
|
||||
}
|
||||
if opts.Provider == "" {
|
||||
opts.Provider = string(enums.KnowledgeChunkProviderStructured)
|
||||
}
|
||||
return opts
|
||||
}
|
||||
|
||||
func normalizeText(text string) string {
|
||||
return strings.Join(strings.Fields(strings.TrimSpace(text)), " ")
|
||||
}
|
||||
|
||||
func estimateTokenCount(text string) int {
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" {
|
||||
return 0
|
||||
}
|
||||
count := 0
|
||||
inWord := false
|
||||
for _, r := range text {
|
||||
switch {
|
||||
case unicode.IsSpace(r):
|
||||
inWord = false
|
||||
case unicode.Is(unicode.Han, r):
|
||||
count++
|
||||
inWord = false
|
||||
case unicode.IsLetter(r) || unicode.IsDigit(r):
|
||||
if !inWord {
|
||||
count++
|
||||
inWord = true
|
||||
}
|
||||
default:
|
||||
count++
|
||||
inWord = false
|
||||
}
|
||||
}
|
||||
if count == 0 {
|
||||
return utf8.RuneCountInString(text)
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func contentHash(text string) string {
|
||||
sum := sha256.Sum256([]byte(text))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func splitSentences(text string) []string {
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
var sentences []string
|
||||
var builder strings.Builder
|
||||
for _, r := range text {
|
||||
builder.WriteRune(r)
|
||||
switch r {
|
||||
case '\n', '。', '!', '?', '!', '?', ';', ';':
|
||||
sentence := normalizeText(builder.String())
|
||||
if sentence != "" {
|
||||
sentences = append(sentences, sentence)
|
||||
}
|
||||
builder.Reset()
|
||||
}
|
||||
}
|
||||
if builder.Len() > 0 {
|
||||
sentence := normalizeText(builder.String())
|
||||
if sentence != "" {
|
||||
sentences = append(sentences, sentence)
|
||||
}
|
||||
}
|
||||
if len(sentences) == 0 {
|
||||
return []string{normalizeText(text)}
|
||||
}
|
||||
return sentences
|
||||
}
|
||||
|
||||
func tailTextByTokens(text string, tokenLimit int) string {
|
||||
if tokenLimit <= 0 {
|
||||
return ""
|
||||
}
|
||||
sentences := splitSentences(text)
|
||||
if len(sentences) == 0 {
|
||||
return ""
|
||||
}
|
||||
var selected []string
|
||||
total := 0
|
||||
for i := len(sentences) - 1; i >= 0; i-- {
|
||||
sentence := sentences[i]
|
||||
tokens := estimateTokenCount(sentence)
|
||||
if total > 0 && total+tokens > tokenLimit {
|
||||
break
|
||||
}
|
||||
selected = append([]string{sentence}, selected...)
|
||||
total += tokens
|
||||
}
|
||||
return strings.TrimSpace(strings.Join(selected, " "))
|
||||
}
|
||||
|
||||
func splitPlainText(text string, opts ChunkOptions) []string {
|
||||
text = normalizeText(text)
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
opts = normalizeOptions(opts)
|
||||
sentences := splitSentences(text)
|
||||
if len(sentences) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
chunks := make([]string, 0)
|
||||
current := make([]string, 0)
|
||||
currentTokens := 0
|
||||
|
||||
flush := func() {
|
||||
if len(current) == 0 {
|
||||
return
|
||||
}
|
||||
chunks = append(chunks, strings.Join(current, " "))
|
||||
}
|
||||
|
||||
for _, sentence := range sentences {
|
||||
sentenceTokens := estimateTokenCount(sentence)
|
||||
if sentenceTokens > opts.MaxTokens {
|
||||
if len(current) > 0 {
|
||||
flush()
|
||||
overlap := tailTextByTokens(strings.Join(current, " "), opts.OverlapTokens)
|
||||
current = nil
|
||||
currentTokens = 0
|
||||
if overlap != "" {
|
||||
current = append(current, overlap)
|
||||
currentTokens = estimateTokenCount(overlap)
|
||||
}
|
||||
}
|
||||
for _, piece := range splitLongSentence(sentence, opts.MaxTokens) {
|
||||
piece = normalizeText(piece)
|
||||
if piece != "" {
|
||||
chunks = append(chunks, piece)
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if currentTokens > 0 && currentTokens+sentenceTokens > opts.MaxTokens {
|
||||
flush()
|
||||
overlap := tailTextByTokens(strings.Join(current, " "), opts.OverlapTokens)
|
||||
current = nil
|
||||
currentTokens = 0
|
||||
if overlap != "" {
|
||||
current = append(current, overlap)
|
||||
currentTokens = estimateTokenCount(overlap)
|
||||
}
|
||||
}
|
||||
|
||||
current = append(current, sentence)
|
||||
currentTokens += sentenceTokens
|
||||
}
|
||||
|
||||
flush()
|
||||
return chunks
|
||||
}
|
||||
|
||||
func splitLongSentence(text string, maxTokens int) []string {
|
||||
runes := []rune(strings.TrimSpace(text))
|
||||
if len(runes) == 0 {
|
||||
return nil
|
||||
}
|
||||
if maxTokens <= 0 {
|
||||
return []string{text}
|
||||
}
|
||||
window := maxTokens * 2
|
||||
if window < 50 {
|
||||
window = 50
|
||||
}
|
||||
var result []string
|
||||
for start := 0; start < len(runes); start += window {
|
||||
end := start + window
|
||||
if end > len(runes) {
|
||||
end = len(runes)
|
||||
}
|
||||
part := normalizeText(string(runes[start:end]))
|
||||
if part != "" {
|
||||
result = append(result, part)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,702 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
ragchunk "cs-agent/internal/ai/rag/chunk"
|
||||
"cs-agent/internal/ai/rag/vectordb"
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/repositories"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/mlogclub/simple/common/strs"
|
||||
"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 := repositories.KnowledgeDocumentRepository.Get(sqls.DB(), documentID)
|
||||
if document == nil {
|
||||
return fmt.Errorf("document not found: %d", documentID)
|
||||
}
|
||||
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 := repositories.KnowledgeBaseRepository.Get(sqls.DB(), document.KnowledgeBaseID)
|
||||
if knowledgeBase == nil {
|
||||
return fail(fmt.Errorf("knowledge base not found: %d", document.KnowledgeBaseID))
|
||||
}
|
||||
|
||||
existingChunks := repositories.KnowledgeChunkRepository.FindByDocumentID(sqls.DB(), document.ID)
|
||||
|
||||
chunks, err := s.registry.Chunk(ctx, &ragchunk.ChunkRequest{
|
||||
KnowledgeBaseID: document.KnowledgeBaseID,
|
||||
DocumentID: document.ID,
|
||||
DocumentTitle: document.Title,
|
||||
ContentType: document.ContentType,
|
||||
Content: document.Content,
|
||||
PlainText: ExtractPlainText(document.Content, document.ContentType),
|
||||
Options: ragchunk.ChunkOptions{
|
||||
Provider: firstNonEmptyString(knowledgeBase.ChunkProvider, s.chunkConfig.Provider),
|
||||
TargetTokens: firstPositiveInt(knowledgeBase.ChunkTargetTokens, s.chunkConfig.TargetTokens),
|
||||
MaxTokens: firstPositiveInt(knowledgeBase.ChunkMaxTokens, s.chunkConfig.MaxTokens),
|
||||
OverlapTokens: firstPositiveInt(knowledgeBase.ChunkOverlapTokens, s.chunkConfig.OverlapTokens),
|
||||
EnableFallback: s.chunkConfig.EnableFallback,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fail(fmt.Errorf("failed to chunk document: %w", err))
|
||||
}
|
||||
if len(chunks) == 0 {
|
||||
return fail(fmt.Errorf("no chunks generated from document"))
|
||||
}
|
||||
|
||||
collectionName := s.getCollectionName()
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fail(fmt.Errorf("vectordb provider not initialized"))
|
||||
}
|
||||
|
||||
if _, err := ai.Embedding.GetModel(ctx); err != nil {
|
||||
return fail(fmt.Errorf("failed to get embedding model: %w", err))
|
||||
}
|
||||
|
||||
existingVectorIDs := make([]string, 0, len(existingChunks))
|
||||
for _, chunk := range existingChunks {
|
||||
if strs.IsNotBlank(chunk.VectorID) {
|
||||
existingVectorIDs = append(existingVectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
|
||||
vectors := make([]vectordb.Vector, 0, len(chunks))
|
||||
chunkModels := make([]models.KnowledgeChunk, 0, len(chunks))
|
||||
dimension := 0
|
||||
|
||||
for i, chunk := range chunks {
|
||||
embeddingResult, err := ai.Embedding.GenerateEmbedding(ctx, chunk.Content)
|
||||
if err != nil {
|
||||
slog.Error("Failed to generate embedding for chunk", "document_id", document.ID, "chunk_index", i, "error", err)
|
||||
return fail(fmt.Errorf("failed to generate embedding for chunk %d: %w", i, err))
|
||||
}
|
||||
if dimension == 0 {
|
||||
dimension = embeddingResult.Dimension
|
||||
}
|
||||
|
||||
chunkID := buildKnowledgeChunkVectorID(knowledgeBase.ID, document.ID, chunk.ChunkNo)
|
||||
providerName := ""
|
||||
if chunk.Metadata != nil {
|
||||
if value, ok := chunk.Metadata["provider"].(string); ok {
|
||||
providerName = value
|
||||
}
|
||||
}
|
||||
chunkModel := models.KnowledgeChunk{
|
||||
KnowledgeBaseID: knowledgeBase.ID,
|
||||
DocumentID: document.ID,
|
||||
ChunkNo: chunk.ChunkNo,
|
||||
Title: chunk.Title,
|
||||
Content: chunk.Content,
|
||||
ContentHash: buildChunkContentHash(chunk.Content),
|
||||
CharCount: chunk.CharCount,
|
||||
TokenCount: chunk.TokenCount,
|
||||
ChunkType: string(chunk.ChunkType),
|
||||
SectionPath: chunk.SectionPath,
|
||||
Provider: providerName,
|
||||
VectorID: chunkID,
|
||||
Status: enums.StatusOk,
|
||||
CreatedAt: time.Now(),
|
||||
UpdatedAt: time.Now(),
|
||||
}
|
||||
chunkModels = append(chunkModels, chunkModel)
|
||||
|
||||
vectors = append(vectors, vectordb.Vector{
|
||||
ID: chunkID,
|
||||
Vector: embeddingResult.Vector,
|
||||
Payload: vectordb.ChunkPayload{
|
||||
KnowledgeBaseID: knowledgeBase.ID,
|
||||
DocumentID: document.ID,
|
||||
DocumentTitle: document.Title,
|
||||
ChunkNo: chunk.ChunkNo,
|
||||
ChunkType: string(chunk.ChunkType),
|
||||
SectionPath: chunk.SectionPath,
|
||||
Content: chunk.Content,
|
||||
Title: chunk.Title,
|
||||
Provider: providerName,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
if len(vectors) == 0 {
|
||||
return fail(fmt.Errorf("no vectors generated"))
|
||||
}
|
||||
|
||||
collectionInfo, err := provider.GetCollection(ctx, collectionName)
|
||||
if err != nil || collectionInfo == nil {
|
||||
if dimension <= 0 {
|
||||
return fail(fmt.Errorf("invalid embedding dimension: %d", dimension))
|
||||
}
|
||||
if err := provider.CreateCollection(ctx, collectionName, dimension); err != nil {
|
||||
return fail(fmt.Errorf("failed to create collection: %w", err))
|
||||
}
|
||||
slog.Info("Created collection for knowledge base", "collection", collectionName, "dimension", dimension)
|
||||
}
|
||||
|
||||
if len(existingVectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, existingVectorIDs); err != nil {
|
||||
return fail(fmt.Errorf("failed to delete old vectors: %w", err))
|
||||
}
|
||||
}
|
||||
|
||||
if err := provider.UpsertVectors(ctx, collectionName, vectors); err != nil {
|
||||
return fail(fmt.Errorf("failed to upsert vectors: %w", err))
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
if err := ctx.Tx.Where("document_id = ?", document.ID).Delete(&models.KnowledgeChunk{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for _, chunk := range chunkModels {
|
||||
if err := ctx.Tx.Create(&chunk).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return fail(fmt.Errorf("failed to save chunks: %w", 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", len(chunks)),
|
||||
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 := repositories.KnowledgeFAQRepository.Get(sqls.DB(), faqID)
|
||||
if faq == nil {
|
||||
return fmt.Errorf("faq not found: %d", faqID)
|
||||
}
|
||||
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 := repositories.KnowledgeBaseRepository.Get(sqls.DB(), faq.KnowledgeBaseID)
|
||||
if knowledgeBase == nil {
|
||||
return fail(fmt.Errorf("knowledge base not found: %d", faq.KnowledgeBaseID))
|
||||
}
|
||||
if knowledgeBase.KnowledgeType != string(enums.KnowledgeBaseTypeFAQ) {
|
||||
return fail(fmt.Errorf("knowledge base %d is not faq type", knowledgeBase.ID))
|
||||
}
|
||||
existingChunks := repositories.KnowledgeChunkRepository.FindByFaqID(sqls.DB(), faq.ID)
|
||||
content := buildFAQChunkContent(faq)
|
||||
if content == "" {
|
||||
return fail(fmt.Errorf("faq content is empty"))
|
||||
}
|
||||
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fail(fmt.Errorf("vectordb provider not initialized"))
|
||||
}
|
||||
if _, err := ai.Embedding.GetModel(ctx); err != nil {
|
||||
return fail(fmt.Errorf("failed to get embedding model: %w", err))
|
||||
}
|
||||
embeddingResult, err := ai.Embedding.GenerateEmbedding(ctx, content)
|
||||
if err != nil {
|
||||
return fail(fmt.Errorf("failed to generate embedding for faq %d: %w", faq.ID, err))
|
||||
}
|
||||
|
||||
chunkID := buildKnowledgeFAQChunkVectorID(knowledgeBase.ID, faq.ID, 0)
|
||||
chunkModel := models.KnowledgeChunk{
|
||||
KnowledgeBaseID: knowledgeBase.ID,
|
||||
FaqID: faq.ID,
|
||||
ChunkNo: 0,
|
||||
Title: faq.Question,
|
||||
Content: content,
|
||||
ContentHash: buildChunkContentHash(content),
|
||||
CharCount: len([]rune(content)),
|
||||
TokenCount: len([]rune(content)) / 2,
|
||||
ChunkType: string(enums.KnowledgeChunkTypeFAQ),
|
||||
Provider: string(enums.KnowledgeChunkProviderFAQ),
|
||||
VectorID: chunkID,
|
||||
Status: enums.StatusOk,
|
||||
CreatedAt: time.Now(),
|
||||
UpdatedAt: time.Now(),
|
||||
}
|
||||
|
||||
collectionName := s.getCollectionName()
|
||||
collectionInfo, err := provider.GetCollection(ctx, collectionName)
|
||||
if err != nil || collectionInfo == nil {
|
||||
if err := provider.CreateCollection(ctx, collectionName, embeddingResult.Dimension); err != nil {
|
||||
return fail(fmt.Errorf("failed to create collection: %w", err))
|
||||
}
|
||||
}
|
||||
|
||||
existingVectorIDs := make([]string, 0, len(existingChunks))
|
||||
for _, chunk := range existingChunks {
|
||||
if strs.IsNotBlank(chunk.VectorID) {
|
||||
existingVectorIDs = append(existingVectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
if len(existingVectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, existingVectorIDs); err != nil {
|
||||
return fail(fmt.Errorf("failed to delete old vectors: %w", err))
|
||||
}
|
||||
}
|
||||
|
||||
if err := provider.UpsertVectors(ctx, collectionName, []vectordb.Vector{{
|
||||
ID: chunkID,
|
||||
Vector: embeddingResult.Vector,
|
||||
Payload: vectordb.ChunkPayload{
|
||||
KnowledgeBaseID: knowledgeBase.ID,
|
||||
FaqID: faq.ID,
|
||||
FaqQuestion: faq.Question,
|
||||
ChunkNo: 0,
|
||||
ChunkType: string(enums.KnowledgeChunkTypeFAQ),
|
||||
Content: content,
|
||||
Title: faq.Question,
|
||||
Provider: string(enums.KnowledgeChunkProviderFAQ),
|
||||
},
|
||||
}}); err != nil {
|
||||
return fail(fmt.Errorf("failed to upsert vectors: %w", err))
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
if err := ctx.Tx.Where("faq_id = ?", faq.ID).Delete(&models.KnowledgeChunk{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return ctx.Tx.Create(&chunkModel).Error
|
||||
}); err != nil {
|
||||
return fail(fmt.Errorf("failed to save faq chunk: %w", 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 {
|
||||
document := repositories.KnowledgeDocumentRepository.Get(sqls.DB(), documentID)
|
||||
if document == nil {
|
||||
return nil
|
||||
}
|
||||
chunks := repositories.KnowledgeChunkRepository.Find(sqls.DB(), sqls.NewCnd().Eq("document_id", documentID))
|
||||
return s.removeDocumentIndexByChunks(ctx, document.KnowledgeBaseID, documentID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) RemoveDocumentIndexFromKnowledgeBase(ctx context.Context, knowledgeBaseID int64, documentID int64) error {
|
||||
chunks := repositories.KnowledgeChunkRepository.Find(sqls.DB(), sqls.NewCnd().Eq("document_id", documentID))
|
||||
return s.removeDocumentIndexByChunks(ctx, knowledgeBaseID, documentID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) RemoveDocumentIndexByChunkModels(ctx context.Context, knowledgeBaseID int64, documentID int64, chunks []models.KnowledgeChunk) error {
|
||||
return s.removeDocumentIndexByChunks(ctx, knowledgeBaseID, documentID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) removeDocumentIndexByChunks(ctx context.Context, knowledgeBaseID int64, documentID int64, chunks []models.KnowledgeChunk) error {
|
||||
if len(chunks) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
collectionName := s.getCollectionName()
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
|
||||
vectorIDs := make([]string, 0, len(chunks))
|
||||
for _, chunk := range chunks {
|
||||
if chunk.VectorID != "" {
|
||||
vectorIDs = append(vectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
|
||||
if len(vectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, vectorIDs); err != nil {
|
||||
slog.Error("Failed to delete vectors", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
return ctx.Tx.Where("document_id = ?", documentID).Delete(&models.KnowledgeChunk{}).Error
|
||||
}); 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 {
|
||||
faq := repositories.KnowledgeFAQRepository.Get(sqls.DB(), faqID)
|
||||
if faq == nil {
|
||||
return nil
|
||||
}
|
||||
chunks := repositories.KnowledgeChunkRepository.FindByFaqID(sqls.DB(), faqID)
|
||||
return s.removeFAQIndexByChunks(ctx, faq.KnowledgeBaseID, faqID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) RemoveFAQIndexByChunkModels(ctx context.Context, knowledgeBaseID int64, faqID int64, chunks []models.KnowledgeChunk) error {
|
||||
return s.removeFAQIndexByChunks(ctx, knowledgeBaseID, faqID, chunks)
|
||||
}
|
||||
|
||||
func (s *index) removeFAQIndexByChunks(ctx context.Context, knowledgeBaseID int64, faqID int64, chunks []models.KnowledgeChunk) error {
|
||||
if len(chunks) == 0 {
|
||||
return nil
|
||||
}
|
||||
collectionName := s.getCollectionName()
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
vectorIDs := make([]string, 0, len(chunks))
|
||||
for _, chunk := range chunks {
|
||||
if chunk.VectorID != "" {
|
||||
vectorIDs = append(vectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
if len(vectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, vectorIDs); err != nil {
|
||||
slog.Error("Failed to delete faq vectors", "error", err)
|
||||
}
|
||||
}
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
return ctx.Tx.Where("faq_id = ?", faqID).Delete(&models.KnowledgeChunk{}).Error
|
||||
}); 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) 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")
|
||||
}
|
||||
|
||||
existing, err := provider.GetCollection(ctx, collectionName)
|
||||
if err == nil && existing != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return provider.CreateCollection(ctx, 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 {
|
||||
if faq == nil {
|
||||
return ""
|
||||
}
|
||||
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) markDocumentIndexPending(documentID int64) error {
|
||||
return repositories.KnowledgeDocumentRepository.Updates(sqls.DB(), documentID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusPending,
|
||||
"indexed_at": nil,
|
||||
"index_error": "",
|
||||
"updated_at": time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markDocumentIndexIndexed(documentID int64) error {
|
||||
now := time.Now()
|
||||
return repositories.KnowledgeDocumentRepository.Updates(sqls.DB(), documentID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusIndexed,
|
||||
"indexed_at": &now,
|
||||
"index_error": "",
|
||||
"updated_at": now,
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markDocumentIndexFailed(documentID int64, err error) error {
|
||||
return repositories.KnowledgeDocumentRepository.Updates(sqls.DB(), documentID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusFailed,
|
||||
"index_error": truncateIndexError(err),
|
||||
"updated_at": time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markKnowledgeBaseDocumentsIndexPending(knowledgeBaseID int64, documentIDs []int64) error {
|
||||
if len(documentIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
return sqls.DB().Model(&models.KnowledgeDocument{}).
|
||||
Where("knowledge_base_id = ?", knowledgeBaseID).
|
||||
Where("id IN ?", documentIDs).
|
||||
Updates(map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusPending,
|
||||
"indexed_at": nil,
|
||||
"index_error": "",
|
||||
"updated_at": time.Now(),
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (s *index) markFAQIndexPending(faqID int64) error {
|
||||
return repositories.KnowledgeFAQRepository.Updates(sqls.DB(), faqID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusPending,
|
||||
"indexed_at": nil,
|
||||
"index_error": "",
|
||||
"updated_at": time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markFAQIndexIndexed(faqID int64) error {
|
||||
now := time.Now()
|
||||
return repositories.KnowledgeFAQRepository.Updates(sqls.DB(), faqID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusIndexed,
|
||||
"indexed_at": &now,
|
||||
"index_error": "",
|
||||
"updated_at": now,
|
||||
})
|
||||
}
|
||||
|
||||
func (s *index) markFAQIndexFailed(faqID int64, err error) error {
|
||||
return repositories.KnowledgeFAQRepository.Updates(sqls.DB(), faqID, map[string]any{
|
||||
"index_status": enums.KnowledgeDocumentIndexStatusFailed,
|
||||
"index_error": truncateIndexError(err),
|
||||
"updated_at": time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
func truncateIndexError(err error) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
message := err.Error()
|
||||
if len(message) <= 1000 {
|
||||
return message
|
||||
}
|
||||
return message[:1000]
|
||||
}
|
||||
|
||||
func (s *index) resetKnowledgeBaseIndexStorage(ctx context.Context, knowledgeBaseID int64) error {
|
||||
collectionName := s.getCollectionName()
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
|
||||
chunks := repositories.KnowledgeChunkRepository.Find(sqls.DB(), sqls.NewCnd().Eq("knowledge_base_id", knowledgeBaseID))
|
||||
vectorIDs := make([]string, 0, len(chunks))
|
||||
for _, chunk := range chunks {
|
||||
if strs.IsNotBlank(chunk.VectorID) {
|
||||
vectorIDs = append(vectorIDs, chunk.VectorID)
|
||||
}
|
||||
}
|
||||
if len(vectorIDs) > 0 {
|
||||
if err := provider.DeleteVectors(ctx, collectionName, vectorIDs); err != nil {
|
||||
return fmt.Errorf("failed to delete vectors for knowledge base %d before rebuild: %w", knowledgeBaseID, err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
return ctx.Tx.Where("knowledge_base_id = ?", knowledgeBaseID).Delete(&models.KnowledgeChunk{}).Error
|
||||
}); err != nil {
|
||||
return fmt.Errorf("failed to clear chunks before rebuild: %w", err)
|
||||
}
|
||||
|
||||
slog.Info("Knowledge base index storage reset",
|
||||
"knowledge_base_id", knowledgeBaseID,
|
||||
"collection", collectionName)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
type rerank struct{}
|
||||
|
||||
var Rerank = &rerank{}
|
||||
|
||||
func (s *rerank) Rerank(ctx context.Context, query string, documents []string, topN int) ([]RerankResult, error) {
|
||||
if len(documents) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if topN <= 0 {
|
||||
topN = len(documents)
|
||||
}
|
||||
|
||||
results, err := s.callRerankAPI(ctx, query, documents, topN)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (s *rerank) callRerankAPI(ctx context.Context, query string, documents []string, topN int) ([]RerankResult, error) {
|
||||
config, err := ai.GetEnabledAIConfig(enums.AIModelTypeRerank)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
reqBody := RerankRequest{
|
||||
Model: config.ModelName,
|
||||
Query: query,
|
||||
Documents: documents,
|
||||
TopN: topN,
|
||||
}
|
||||
|
||||
jsonBody, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", config.BaseURL+"/v1/rerank", bytes.NewBuffer(jsonBody))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+config.APIKey)
|
||||
|
||||
client := &http.Client{
|
||||
Timeout: time.Duration(config.TimeoutMS) * time.Millisecond,
|
||||
}
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to call rerank API: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||
}
|
||||
|
||||
var rerankResp RerankResponse
|
||||
if err := json.Unmarshal(body, &rerankResp); err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
results := make([]RerankResult, 0, len(rerankResp.Results))
|
||||
for _, r := range rerankResp.Results {
|
||||
results = append(results, RerankResult{
|
||||
Index: r.Index,
|
||||
RelevanceScore: r.RelevanceScore,
|
||||
})
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (s *rerank) RerankResults(ctx context.Context, query string, results []RetrieveResult, topN int) ([]RetrieveResult, error) {
|
||||
if len(results) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if topN <= 0 {
|
||||
topN = len(results)
|
||||
}
|
||||
|
||||
documents := make([]string, 0, len(results))
|
||||
for _, r := range results {
|
||||
documents = append(documents, r.Content)
|
||||
}
|
||||
|
||||
rerankResults, err := s.Rerank(ctx, query, documents, topN)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
rerankedResults := make([]RetrieveResult, 0, len(rerankResults))
|
||||
for _, rr := range rerankResults {
|
||||
if rr.Index < len(results) {
|
||||
result := results[rr.Index]
|
||||
result.Score = float32(rr.RelevanceScore)
|
||||
rerankedResults = append(rerankedResults, result)
|
||||
}
|
||||
}
|
||||
|
||||
return rerankedResults, nil
|
||||
}
|
||||
|
||||
func (s *rerank) SimpleRerank(query string, results []RetrieveResult, topN int) []RetrieveResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
sort.Slice(results, func(i, j int) bool {
|
||||
return results[i].Score > results[j].Score
|
||||
})
|
||||
|
||||
if topN > 0 && len(results) > topN {
|
||||
return results[:topN]
|
||||
}
|
||||
|
||||
return results
|
||||
}
|
||||
@@ -0,0 +1,507 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
|
||||
"cs-agent/internal/ai"
|
||||
"cs-agent/internal/ai/rag/vectordb"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"cs-agent/internal/repositories"
|
||||
)
|
||||
|
||||
type retrieve struct {
|
||||
}
|
||||
|
||||
var Retrieve = &retrieve{}
|
||||
|
||||
func (s *retrieve) Retrieve(ctx context.Context, req RetrieveRequest) ([]RetrieveResult, error) {
|
||||
results, _, err := s.RetrieveWithTrace(ctx, req)
|
||||
return results, err
|
||||
}
|
||||
|
||||
type RetrieveTrace struct {
|
||||
EmbeddingMs int64
|
||||
VectorSearchMs int64
|
||||
HydrateMs int64
|
||||
}
|
||||
|
||||
func (s *retrieve) RetrieveWithTrace(ctx context.Context, req RetrieveRequest) ([]RetrieveResult, *RetrieveTrace, error) {
|
||||
trace := &RetrieveTrace{}
|
||||
if req.Query == "" {
|
||||
return nil, trace, nil
|
||||
}
|
||||
knowledgeBaseIDs := normalizeKnowledgeBaseIDs(req.KnowledgeBaseIDs)
|
||||
if len(knowledgeBaseIDs) == 0 {
|
||||
return nil, trace, nil
|
||||
}
|
||||
|
||||
retrievableKnowledgeBases := s.loadRetrievableKnowledgeBases(knowledgeBaseIDs)
|
||||
if len(retrievableKnowledgeBases) == 0 {
|
||||
slog.Info("Skip retrieve for non-enabled knowledge bases",
|
||||
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs))
|
||||
return nil, trace, nil
|
||||
}
|
||||
|
||||
embeddingStartedAt := time.Now()
|
||||
embeddingResult, err := ai.Embedding.GenerateEmbedding(ctx, req.Query)
|
||||
trace.EmbeddingMs = time.Since(embeddingStartedAt).Milliseconds()
|
||||
if err != nil {
|
||||
return nil, trace, fmt.Errorf("failed to generate query embedding: %w", err)
|
||||
}
|
||||
|
||||
collectionName := knowledgeCollectionName
|
||||
provider := vectordb.GetProvider()
|
||||
if provider == nil {
|
||||
return nil, trace, fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
|
||||
searchResults := make([]vectordb.SearchResult, 0)
|
||||
vectorSearchStartedAt := time.Now()
|
||||
for _, knowledgeBase := range retrievableKnowledgeBases {
|
||||
topK, scoreThreshold := resolveKnowledgeBaseSearchOptions(req, &knowledgeBase)
|
||||
kbResults, searchErr := provider.Search(ctx, &vectordb.SearchRequest{
|
||||
CollectionName: collectionName,
|
||||
Vector: embeddingResult.Vector,
|
||||
TopK: topK,
|
||||
ScoreThreshold: scoreThreshold,
|
||||
Filter: &vectordb.SearchFilter{
|
||||
KnowledgeBaseIDs: []int64{knowledgeBase.ID},
|
||||
},
|
||||
})
|
||||
if searchErr != nil {
|
||||
slog.Error("Failed to search vectors",
|
||||
"knowledge_base_id", knowledgeBase.ID,
|
||||
"error", searchErr)
|
||||
trace.VectorSearchMs = time.Since(vectorSearchStartedAt).Milliseconds()
|
||||
return nil, trace, fmt.Errorf("failed to search vectors: %w", searchErr)
|
||||
}
|
||||
if len(kbResults) == 0 && scoreThreshold > 0 {
|
||||
s.logEmptySearchDiagnostics(ctx, provider, collectionName, embeddingResult.Vector, topK, scoreThreshold, []int64{knowledgeBase.ID}, req)
|
||||
}
|
||||
searchResults = append(searchResults, kbResults...)
|
||||
}
|
||||
trace.VectorSearchMs = time.Since(vectorSearchStartedAt).Milliseconds()
|
||||
|
||||
if len(searchResults) == 0 {
|
||||
return nil, trace, nil
|
||||
}
|
||||
sort.SliceStable(searchResults, func(i, j int) bool {
|
||||
if searchResults[i].Score == searchResults[j].Score {
|
||||
return searchResults[i].ID < searchResults[j].ID
|
||||
}
|
||||
return searchResults[i].Score > searchResults[j].Score
|
||||
})
|
||||
|
||||
results := make([]RetrieveResult, 0, len(searchResults))
|
||||
hydrateStartedAt := time.Now()
|
||||
vectorIDs := make([]string, 0, len(searchResults))
|
||||
for _, sr := range searchResults {
|
||||
if strings.TrimSpace(sr.ID) == "" {
|
||||
continue
|
||||
}
|
||||
vectorIDs = append(vectorIDs, sr.ID)
|
||||
}
|
||||
chunks := repositories.KnowledgeChunkRepository.FindByVectorIDs(sqls.DB(), vectorIDs)
|
||||
chunkByVectorID := make(map[string]*models.KnowledgeChunk, len(chunks))
|
||||
documentIDs := make([]int64, 0)
|
||||
faqIDs := make([]int64, 0)
|
||||
documentSeen := make(map[int64]struct{})
|
||||
faqSeen := make(map[int64]struct{})
|
||||
for i := range chunks {
|
||||
chunk := &chunks[i]
|
||||
chunkByVectorID[chunk.VectorID] = chunk
|
||||
if chunk.DocumentID > 0 {
|
||||
if _, ok := documentSeen[chunk.DocumentID]; !ok {
|
||||
documentSeen[chunk.DocumentID] = struct{}{}
|
||||
documentIDs = append(documentIDs, chunk.DocumentID)
|
||||
}
|
||||
}
|
||||
if chunk.FaqID > 0 {
|
||||
if _, ok := faqSeen[chunk.FaqID]; !ok {
|
||||
faqSeen[chunk.FaqID] = struct{}{}
|
||||
faqIDs = append(faqIDs, chunk.FaqID)
|
||||
}
|
||||
}
|
||||
}
|
||||
documents := repositories.KnowledgeDocumentRepository.FindByIDs(sqls.DB(), documentIDs)
|
||||
documentByID := make(map[int64]*models.KnowledgeDocument, len(documents))
|
||||
for i := range documents {
|
||||
document := &documents[i]
|
||||
documentByID[document.ID] = document
|
||||
}
|
||||
faqs := repositories.KnowledgeFAQRepository.FindByIDs(sqls.DB(), faqIDs)
|
||||
faqByID := make(map[int64]*models.KnowledgeFAQ, len(faqs))
|
||||
for i := range faqs {
|
||||
faq := &faqs[i]
|
||||
faqByID[faq.ID] = faq
|
||||
}
|
||||
for _, sr := range searchResults {
|
||||
chunk := chunkByVectorID[sr.ID]
|
||||
if chunk == nil || chunk.Status != enums.StatusOk {
|
||||
continue
|
||||
}
|
||||
|
||||
documentTitle := ""
|
||||
faqQuestion := ""
|
||||
if chunk.DocumentID > 0 {
|
||||
document := documentByID[chunk.DocumentID]
|
||||
if document == nil || document.Status != enums.StatusOk {
|
||||
continue
|
||||
}
|
||||
documentTitle = document.Title
|
||||
}
|
||||
if chunk.FaqID > 0 {
|
||||
faq := faqByID[chunk.FaqID]
|
||||
if faq == nil || faq.Status != enums.StatusOk {
|
||||
continue
|
||||
}
|
||||
faqQuestion = faq.Question
|
||||
}
|
||||
|
||||
results = append(results, RetrieveResult{
|
||||
KnowledgeBaseID: chunk.KnowledgeBaseID,
|
||||
ChunkID: chunk.ID,
|
||||
DocumentID: chunk.DocumentID,
|
||||
DocumentTitle: documentTitle,
|
||||
FaqID: chunk.FaqID,
|
||||
FaqQuestion: faqQuestion,
|
||||
ChunkNo: chunk.ChunkNo,
|
||||
Title: chunk.Title,
|
||||
SectionPath: chunk.SectionPath,
|
||||
Content: chunk.Content,
|
||||
Score: sr.Score,
|
||||
ChunkType: extractChunkType(sr.Payload),
|
||||
})
|
||||
}
|
||||
trace.HydrateMs = time.Since(hydrateStartedAt).Milliseconds()
|
||||
|
||||
return results, trace, nil
|
||||
}
|
||||
|
||||
func extractChunkType(payload vectordb.ChunkPayload) string {
|
||||
if payload.ChunkType != "" {
|
||||
return payload.ChunkType
|
||||
}
|
||||
return string(enums.KnowledgeChunkTypeText)
|
||||
}
|
||||
|
||||
func (s *retrieve) logEmptySearchDiagnostics(ctx context.Context, provider vectordb.Provider, collectionName string, vector []float32, topK int, scoreThreshold float32, knowledgeBaseIDs []int64, req RetrieveRequest) {
|
||||
rawResults, err := provider.Search(ctx, &vectordb.SearchRequest{
|
||||
CollectionName: collectionName,
|
||||
Vector: vector,
|
||||
TopK: topK,
|
||||
ScoreThreshold: 0,
|
||||
Filter: &vectordb.SearchFilter{
|
||||
KnowledgeBaseIDs: knowledgeBaseIDs,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
slog.Warn("Knowledge retrieve diagnostics failed",
|
||||
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs),
|
||||
"collection", collectionName,
|
||||
"query", truncateForLog(req.Query, 80),
|
||||
"score_threshold", scoreThreshold,
|
||||
"error", err)
|
||||
return
|
||||
}
|
||||
if len(rawResults) == 0 {
|
||||
slog.Info("Knowledge retrieve returned no candidates even without threshold",
|
||||
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs),
|
||||
"collection", collectionName,
|
||||
"query", truncateForLog(req.Query, 80),
|
||||
"score_threshold", scoreThreshold)
|
||||
return
|
||||
}
|
||||
|
||||
candidates := make([]string, 0, len(rawResults))
|
||||
for _, item := range rawResults {
|
||||
candidates = append(candidates, fmt.Sprintf("%s:%.4f", item.ID, item.Score))
|
||||
}
|
||||
|
||||
slog.Info("Knowledge retrieve filtered all candidates by score threshold",
|
||||
"knowledge_base_ids", fmt.Sprint(knowledgeBaseIDs),
|
||||
"collection", collectionName,
|
||||
"query", truncateForLog(req.Query, 80),
|
||||
"score_threshold", scoreThreshold,
|
||||
"top_candidates", strings.Join(candidates, ","))
|
||||
}
|
||||
|
||||
func truncateForLog(text string, limit int) string {
|
||||
if limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(strings.TrimSpace(text))
|
||||
if len(runes) <= limit {
|
||||
return string(runes)
|
||||
}
|
||||
return string(runes[:limit]) + "..."
|
||||
}
|
||||
|
||||
func (s *retrieve) RetrieveWithRerank(ctx context.Context, req RetrieveRequest, rerankLimit int) ([]RetrieveResult, error) {
|
||||
results, err := s.Retrieve(ctx, req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if len(results) <= rerankLimit {
|
||||
return results, nil
|
||||
}
|
||||
|
||||
rerankedResults, err := s.rerank(ctx, req.Query, results, rerankLimit)
|
||||
if err != nil {
|
||||
slog.Warn("Rerank failed, returning original results", "error", err)
|
||||
if len(results) > rerankLimit {
|
||||
return results[:rerankLimit], nil
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
return rerankedResults, nil
|
||||
}
|
||||
|
||||
func (s *retrieve) rerank(ctx context.Context, query string, results []RetrieveResult, limit int) ([]RetrieveResult, error) {
|
||||
return Rerank.RerankResults(ctx, query, results, limit)
|
||||
}
|
||||
|
||||
func (s *retrieve) SelectContextResults(results []RetrieveResult, maxTokens int) []RetrieveResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
normalizedResults := normalizeContextResults(results)
|
||||
selected := make([]RetrieveResult, 0, len(normalizedResults))
|
||||
totalTokens := 0
|
||||
documentUsage := make(map[int64]int)
|
||||
|
||||
for _, item := range normalizedResults {
|
||||
if documentUsage[item.DocumentID] >= 2 {
|
||||
continue
|
||||
}
|
||||
chunkText := buildContextChunkText(item)
|
||||
estimatedTokens := len(chunkText) / 2
|
||||
if totalTokens+estimatedTokens > maxTokens {
|
||||
break
|
||||
}
|
||||
selected = append(selected, item)
|
||||
totalTokens += estimatedTokens
|
||||
documentUsage[item.DocumentID]++
|
||||
}
|
||||
return selected
|
||||
}
|
||||
|
||||
func (s *retrieve) BuildContext(ctx context.Context, results []RetrieveResult, maxTokens int) string {
|
||||
if len(results) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
normalizedResults := s.SelectContextResults(results, maxTokens)
|
||||
context := ""
|
||||
for _, r := range normalizedResults {
|
||||
chunkText := buildContextChunkText(r)
|
||||
context += chunkText
|
||||
}
|
||||
|
||||
return context
|
||||
}
|
||||
|
||||
func normalizeContextResults(results []RetrieveResult) []RetrieveResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
merged := mergeAdjacentResults(results)
|
||||
return dedupeSectionResults(merged)
|
||||
}
|
||||
|
||||
func dedupeSectionResults(results []RetrieveResult) []RetrieveResult {
|
||||
seen := make(map[string]struct{})
|
||||
deduped := make([]RetrieveResult, 0, len(results))
|
||||
for _, item := range results {
|
||||
key := buildSectionKey(item)
|
||||
if _, ok := seen[key]; ok {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
deduped = append(deduped, item)
|
||||
}
|
||||
return deduped
|
||||
}
|
||||
|
||||
func mergeAdjacentResults(results []RetrieveResult) []RetrieveResult {
|
||||
if len(results) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
merged := make([]RetrieveResult, 0, len(results))
|
||||
for _, item := range results {
|
||||
if len(merged) == 0 {
|
||||
merged = append(merged, item)
|
||||
continue
|
||||
}
|
||||
|
||||
last := &merged[len(merged)-1]
|
||||
if canMergeContextResult(*last, item) {
|
||||
last.Content = strings.TrimSpace(last.Content + "\n" + item.Content)
|
||||
if item.Score > last.Score {
|
||||
last.Score = item.Score
|
||||
}
|
||||
continue
|
||||
}
|
||||
merged = append(merged, item)
|
||||
}
|
||||
return merged
|
||||
}
|
||||
|
||||
func canMergeContextResult(left, right RetrieveResult) bool {
|
||||
if left.FaqID > 0 || right.FaqID > 0 {
|
||||
return false
|
||||
}
|
||||
if left.DocumentID != right.DocumentID {
|
||||
return false
|
||||
}
|
||||
if left.SectionPath == "" || right.SectionPath == "" {
|
||||
return false
|
||||
}
|
||||
if left.SectionPath != right.SectionPath {
|
||||
return false
|
||||
}
|
||||
return right.ChunkNo == left.ChunkNo+1
|
||||
}
|
||||
|
||||
func buildSectionKey(item RetrieveResult) string {
|
||||
if item.FaqID > 0 {
|
||||
return fmt.Sprintf("faq:%d", item.FaqID)
|
||||
}
|
||||
sectionPath := strings.TrimSpace(item.SectionPath)
|
||||
if sectionPath != "" {
|
||||
return fmt.Sprintf("%d|%s", item.DocumentID, sectionPath)
|
||||
}
|
||||
title := strings.TrimSpace(item.Title)
|
||||
if title != "" {
|
||||
return fmt.Sprintf("%d|%s", item.DocumentID, title)
|
||||
}
|
||||
return fmt.Sprintf("%d|chunk:%d", item.DocumentID, item.ChunkNo)
|
||||
}
|
||||
|
||||
func buildContextChunkText(item RetrieveResult) string {
|
||||
if item.FaqID > 0 {
|
||||
title := strings.TrimSpace(item.FaqQuestion)
|
||||
if title == "" {
|
||||
title = strings.TrimSpace(item.Title)
|
||||
}
|
||||
if title == "" {
|
||||
title = fmt.Sprintf("FAQ#%d", item.FaqID)
|
||||
}
|
||||
return fmt.Sprintf("【FAQ:%s】\n%s\n\n", title, item.Content)
|
||||
}
|
||||
title := strings.TrimSpace(item.DocumentTitle)
|
||||
if title == "" {
|
||||
title = fmt.Sprintf("文档#%d", item.DocumentID)
|
||||
}
|
||||
if item.SectionPath != "" {
|
||||
return fmt.Sprintf("【文档:%s|章节:%s】\n%s\n\n", title, item.SectionPath, item.Content)
|
||||
}
|
||||
if item.Title != "" {
|
||||
return fmt.Sprintf("【文档:%s|标题:%s】\n%s\n\n", title, item.Title, item.Content)
|
||||
}
|
||||
return fmt.Sprintf("【文档:%s】\n%s\n\n", title, item.Content)
|
||||
}
|
||||
|
||||
func (s *retrieve) GetKnowledgeBaseStats(ctx context.Context, knowledgeBaseID int64) (*KnowledgeBaseStats, error) {
|
||||
knowledgeBase := repositories.KnowledgeBaseRepository.Get(sqls.DB(), knowledgeBaseID)
|
||||
if knowledgeBase == nil {
|
||||
return nil, fmt.Errorf("knowledge base not found")
|
||||
}
|
||||
|
||||
documentCount := repositories.KnowledgeDocumentRepository.CountByKnowledgeBaseID(sqls.DB(), knowledgeBaseID)
|
||||
chunkCount := repositories.KnowledgeChunkRepository.CountByKnowledgeBaseID(sqls.DB(), knowledgeBaseID)
|
||||
|
||||
publishedCount := repositories.KnowledgeDocumentRepository.Count(sqls.DB(), sqls.NewCnd().
|
||||
Eq("knowledge_base_id", knowledgeBaseID).
|
||||
Eq("status", enums.StatusOk))
|
||||
|
||||
return &KnowledgeBaseStats{
|
||||
KnowledgeBaseID: knowledgeBaseID,
|
||||
DocumentCount: documentCount,
|
||||
PublishedCount: publishedCount,
|
||||
ChunkCount: chunkCount,
|
||||
VectorCount: int(chunkCount),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeKnowledgeBaseIDs(ids []int64) []int64 {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
seen := make(map[int64]struct{}, len(ids))
|
||||
normalized := make([]int64, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
normalized = append(normalized, id)
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func resolveKnowledgeBaseSearchOptions(req RetrieveRequest, knowledgeBase *models.KnowledgeBase) (int, float32) {
|
||||
topK := req.TopK
|
||||
if topK <= 0 && knowledgeBase != nil && knowledgeBase.DefaultTopK > 0 {
|
||||
topK = knowledgeBase.DefaultTopK
|
||||
}
|
||||
if topK <= 0 {
|
||||
topK = 8
|
||||
}
|
||||
|
||||
scoreThreshold := float32(req.ScoreThreshold)
|
||||
if scoreThreshold <= 0 && knowledgeBase != nil && knowledgeBase.DefaultScoreThreshold > 0 {
|
||||
scoreThreshold = float32(knowledgeBase.DefaultScoreThreshold)
|
||||
}
|
||||
if scoreThreshold <= 0 {
|
||||
scoreThreshold = 0.3
|
||||
}
|
||||
return topK, scoreThreshold
|
||||
}
|
||||
|
||||
func (s *retrieve) loadRetrievableKnowledgeBases(ids []int64) []models.KnowledgeBase {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
items := repositories.KnowledgeBaseRepository.Find(sqls.DB(), sqls.NewCnd().In("id", ids))
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
allowed := make(map[int64]models.KnowledgeBase, len(items))
|
||||
for _, item := range items {
|
||||
if item.Status == enums.StatusOk {
|
||||
allowed[item.ID] = item
|
||||
}
|
||||
}
|
||||
filtered := make([]models.KnowledgeBase, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if item, ok := allowed[id]; ok {
|
||||
filtered = append(filtered, item)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
type KnowledgeBaseStats struct {
|
||||
KnowledgeBaseID int64 `json:"knowledgeBaseId"`
|
||||
DocumentCount int64 `json:"documentCount"`
|
||||
PublishedCount int64 `json:"publishedCount"`
|
||||
ChunkCount int64 `json:"chunkCount"`
|
||||
VectorCount int `json:"vectorCount"`
|
||||
}
|
||||
@@ -0,0 +1,300 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
"cs-agent/internal/pkg/dto"
|
||||
"cs-agent/internal/pkg/dto/response"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
)
|
||||
|
||||
var RetrieveLog = &retrieveLog{}
|
||||
|
||||
type retrieveLog struct {
|
||||
}
|
||||
|
||||
type CreateRetrieveLogRequest struct {
|
||||
KnowledgeBaseID int64
|
||||
Channel string
|
||||
Scene string
|
||||
SessionID string
|
||||
ConversationID int64
|
||||
Question string
|
||||
RewriteQuestion string
|
||||
Answer string
|
||||
AnswerStatus int
|
||||
ChunkProvider string
|
||||
ChunkTargetTokens int
|
||||
ChunkMaxTokens int
|
||||
ChunkOverlapTokens int
|
||||
RerankEnabled bool
|
||||
RerankLimit int
|
||||
Hits []response.KnowledgeSearchResult
|
||||
UsedHits []response.KnowledgeSearchResult
|
||||
Citations []response.KnowledgeCitation
|
||||
LatencyMs int64
|
||||
RetrieveMs int64
|
||||
GenerateMs int64
|
||||
PromptTokens int
|
||||
CompletionTokens int
|
||||
ModelName string
|
||||
}
|
||||
|
||||
type retrieveTraceData struct {
|
||||
Retrieve retrieveTraceRetrieve `json:"retrieve"`
|
||||
ChunkConfig retrieveTraceChunkConfig `json:"chunkConfig"`
|
||||
Context retrieveTraceContext `json:"context"`
|
||||
Citations []retrieveTraceCitation `json:"citations"`
|
||||
}
|
||||
|
||||
type retrieveTraceRetrieve struct {
|
||||
Provider string `json:"provider"`
|
||||
RerankEnabled bool `json:"rerankEnabled"`
|
||||
RerankLimit int `json:"rerankLimit"`
|
||||
RawHitCount int `json:"rawHitCount"`
|
||||
ContextHitCount int `json:"contextHitCount"`
|
||||
CitationCount int `json:"citationCount"`
|
||||
}
|
||||
|
||||
type retrieveTraceChunkConfig struct {
|
||||
Provider string `json:"provider"`
|
||||
TargetTokens int `json:"targetTokens"`
|
||||
MaxTokens int `json:"maxTokens"`
|
||||
OverlapTokens int `json:"overlapTokens"`
|
||||
}
|
||||
|
||||
type retrieveTraceContext struct {
|
||||
KnowledgeBaseIDs []int64 `json:"knowledgeBaseIds"`
|
||||
DocumentIDs []int64 `json:"documentIds"`
|
||||
SectionPaths []string `json:"sectionPaths"`
|
||||
UsedChunkKeys []string `json:"usedChunkKeys"`
|
||||
}
|
||||
|
||||
type retrieveTraceCitation struct {
|
||||
DocumentID int64 `json:"documentId"`
|
||||
ChunkNo int `json:"chunkNo"`
|
||||
SectionPath string `json:"sectionPath"`
|
||||
}
|
||||
|
||||
func (s *retrieveLog) FindHitsByRetrieveLogID(retrieveLogID int64) []models.KnowledgeRetrieveHit {
|
||||
if retrieveLogID <= 0 {
|
||||
return nil
|
||||
}
|
||||
var list []models.KnowledgeRetrieveHit
|
||||
sqls.DB().Where("retrieve_log_id = ?", retrieveLogID).Order("rank_no asc, id asc").Find(&list)
|
||||
return list
|
||||
}
|
||||
|
||||
func (s *retrieveLog) CreateRetrieveLog(req *CreateRetrieveLogRequest, _ *dto.AuthPrincipal) (*models.KnowledgeRetrieveLog, error) {
|
||||
if req == nil {
|
||||
return nil, fmt.Errorf("retrieve log request is nil")
|
||||
}
|
||||
now := time.Now()
|
||||
topScore := 0.0
|
||||
if len(req.Hits) > 0 {
|
||||
topScore = req.Hits[0].Score
|
||||
}
|
||||
traceData := buildRetrieveTraceData(req)
|
||||
|
||||
log := &models.KnowledgeRetrieveLog{
|
||||
KnowledgeBaseID: req.KnowledgeBaseID,
|
||||
Channel: req.Channel,
|
||||
Scene: req.Scene,
|
||||
SessionID: req.SessionID,
|
||||
ConversationID: req.ConversationID,
|
||||
RequestID: uuid.New().String(),
|
||||
Question: req.Question,
|
||||
RewriteQuestion: req.RewriteQuestion,
|
||||
Answer: req.Answer,
|
||||
AnswerStatus: req.AnswerStatus,
|
||||
HitCount: len(req.Hits),
|
||||
TopScore: topScore,
|
||||
ChunkProvider: req.ChunkProvider,
|
||||
ChunkTargetTokens: req.ChunkTargetTokens,
|
||||
ChunkMaxTokens: req.ChunkMaxTokens,
|
||||
ChunkOverlapTokens: req.ChunkOverlapTokens,
|
||||
RerankEnabled: req.RerankEnabled,
|
||||
RerankLimit: req.RerankLimit,
|
||||
CitationCount: len(req.Citations),
|
||||
UsedChunkCount: len(req.UsedHits),
|
||||
LatencyMs: req.LatencyMs,
|
||||
RetrieveMs: req.RetrieveMs,
|
||||
GenerateMs: req.GenerateMs,
|
||||
PromptTokens: req.PromptTokens,
|
||||
CompletionTokens: req.CompletionTokens,
|
||||
ModelName: req.ModelName,
|
||||
TraceData: traceData,
|
||||
CreatedAt: now,
|
||||
}
|
||||
|
||||
usedHitKeys := make(map[string]struct{}, len(req.UsedHits))
|
||||
for _, item := range req.UsedHits {
|
||||
usedHitKeys[buildKnowledgeSearchResultKey(item)] = struct{}{}
|
||||
}
|
||||
citationKeys := make(map[string]struct{}, len(req.Citations))
|
||||
for _, item := range req.Citations {
|
||||
citationKeys[buildKnowledgeCitationKey(item)] = struct{}{}
|
||||
}
|
||||
|
||||
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
if err := ctx.Tx.Create(log).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
for i, hit := range req.Hits {
|
||||
hitKey := buildKnowledgeSearchResultKey(hit)
|
||||
hitRecord := &models.KnowledgeRetrieveHit{
|
||||
RetrieveLogID: log.ID,
|
||||
KnowledgeBaseID: hit.KnowledgeBaseID,
|
||||
ChunkID: hit.ChunkID,
|
||||
DocumentID: hit.DocumentID,
|
||||
DocumentTitle: hit.DocumentTitle,
|
||||
FaqID: hit.FaqID,
|
||||
FaqQuestion: hit.FaqQuestion,
|
||||
ChunkNo: hit.ChunkNo,
|
||||
Title: hit.Title,
|
||||
SectionPath: hit.SectionPath,
|
||||
ChunkType: "",
|
||||
Provider: req.ChunkProvider,
|
||||
RankNo: i + 1,
|
||||
Score: hit.Score,
|
||||
RerankScore: hit.RerankScore,
|
||||
UsedInAnswer: hasHitKey(usedHitKeys, hitKey),
|
||||
IsCitation: hasHitKey(citationKeys, hitKey),
|
||||
Snippet: hit.Content,
|
||||
CreatedAt: now,
|
||||
}
|
||||
if err := ctx.Tx.Create(hitRecord).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return log, nil
|
||||
}
|
||||
|
||||
func buildRetrieveTraceData(req *CreateRetrieveLogRequest) string {
|
||||
trace := retrieveTraceData{
|
||||
Retrieve: retrieveTraceRetrieve{
|
||||
Provider: req.ChunkProvider,
|
||||
RerankEnabled: req.RerankEnabled,
|
||||
RerankLimit: req.RerankLimit,
|
||||
RawHitCount: len(req.Hits),
|
||||
ContextHitCount: len(req.UsedHits),
|
||||
CitationCount: len(req.Citations),
|
||||
},
|
||||
ChunkConfig: retrieveTraceChunkConfig{
|
||||
Provider: req.ChunkProvider,
|
||||
TargetTokens: req.ChunkTargetTokens,
|
||||
MaxTokens: req.ChunkMaxTokens,
|
||||
OverlapTokens: req.ChunkOverlapTokens,
|
||||
},
|
||||
Context: retrieveTraceContext{
|
||||
KnowledgeBaseIDs: distinctKnowledgeBaseIDs(req.UsedHits),
|
||||
DocumentIDs: distinctDocumentIDs(req.UsedHits),
|
||||
SectionPaths: distinctSectionPaths(req.UsedHits),
|
||||
UsedChunkKeys: buildUsedChunkKeys(req.UsedHits),
|
||||
},
|
||||
Citations: buildTraceCitations(req.Citations),
|
||||
}
|
||||
data, err := json.Marshal(trace)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(data)
|
||||
}
|
||||
|
||||
func buildTraceCitations(citations []response.KnowledgeCitation) []retrieveTraceCitation {
|
||||
items := make([]retrieveTraceCitation, 0, len(citations))
|
||||
for _, item := range citations {
|
||||
items = append(items, retrieveTraceCitation{
|
||||
DocumentID: item.DocumentID,
|
||||
ChunkNo: item.ChunkNo,
|
||||
SectionPath: item.SectionPath,
|
||||
})
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func buildUsedChunkKeys(hits []response.KnowledgeSearchResult) []string {
|
||||
keys := make([]string, 0, len(hits))
|
||||
for _, item := range hits {
|
||||
keys = append(keys, buildKnowledgeSearchResultKey(item))
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func distinctKnowledgeBaseIDs(hits []response.KnowledgeSearchResult) []int64 {
|
||||
ids := make([]int64, 0)
|
||||
seen := make(map[int64]struct{})
|
||||
for _, item := range hits {
|
||||
if item.KnowledgeBaseID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[item.KnowledgeBaseID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item.KnowledgeBaseID] = struct{}{}
|
||||
ids = append(ids, item.KnowledgeBaseID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func distinctDocumentIDs(hits []response.KnowledgeSearchResult) []int64 {
|
||||
seen := make(map[int64]struct{})
|
||||
items := make([]int64, 0)
|
||||
for _, item := range hits {
|
||||
if item.DocumentID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[item.DocumentID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item.DocumentID] = struct{}{}
|
||||
items = append(items, item.DocumentID)
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func distinctSectionPaths(hits []response.KnowledgeSearchResult) []string {
|
||||
seen := make(map[string]struct{})
|
||||
items := make([]string, 0)
|
||||
for _, item := range hits {
|
||||
sectionPath := item.SectionPath
|
||||
if sectionPath == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[sectionPath]; ok {
|
||||
continue
|
||||
}
|
||||
seen[sectionPath] = struct{}{}
|
||||
items = append(items, sectionPath)
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func buildKnowledgeSearchResultKey(item response.KnowledgeSearchResult) string {
|
||||
if item.FaqID > 0 {
|
||||
return fmt.Sprintf("faq:%d|%d", item.FaqID, item.ChunkNo)
|
||||
}
|
||||
return fmt.Sprintf("%d|%s|%d", item.DocumentID, item.SectionPath, item.ChunkNo)
|
||||
}
|
||||
|
||||
func buildKnowledgeCitationKey(item response.KnowledgeCitation) string {
|
||||
if item.FaqID > 0 {
|
||||
return fmt.Sprintf("faq:%d|%d", item.FaqID, item.ChunkNo)
|
||||
}
|
||||
return fmt.Sprintf("%d|%s|%d", item.DocumentID, item.SectionPath, item.ChunkNo)
|
||||
}
|
||||
|
||||
func hasHitKey(items map[string]struct{}, key string) bool {
|
||||
_, ok := items[key]
|
||||
return ok
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"cs-agent/internal/models"
|
||||
)
|
||||
|
||||
func TestResolveKnowledgeBaseSearchOptionsUsesKnowledgeBaseDefaults(t *testing.T) {
|
||||
topK, scoreThreshold := resolveKnowledgeBaseSearchOptions(RetrieveRequest{}, &models.KnowledgeBase{
|
||||
DefaultTopK: 6,
|
||||
DefaultScoreThreshold: 0.42,
|
||||
})
|
||||
|
||||
if topK != 6 {
|
||||
t.Fatalf("expected topK 6, got %d", topK)
|
||||
}
|
||||
if scoreThreshold != float32(0.42) {
|
||||
t.Fatalf("expected score threshold 0.42, got %v", scoreThreshold)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveKnowledgeBaseSearchOptionsRequestOverridesKnowledgeBaseDefaults(t *testing.T) {
|
||||
topK, scoreThreshold := resolveKnowledgeBaseSearchOptions(RetrieveRequest{
|
||||
TopK: 9,
|
||||
ScoreThreshold: 0.55,
|
||||
}, &models.KnowledgeBase{
|
||||
DefaultTopK: 6,
|
||||
DefaultScoreThreshold: 0.42,
|
||||
})
|
||||
|
||||
if topK != 9 {
|
||||
t.Fatalf("expected request topK 9, got %d", topK)
|
||||
}
|
||||
if scoreThreshold != float32(0.55) {
|
||||
t.Fatalf("expected request score threshold 0.55, got %v", scoreThreshold)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveKnowledgeBaseSearchOptionsUsesSystemDefaults(t *testing.T) {
|
||||
topK, scoreThreshold := resolveKnowledgeBaseSearchOptions(RetrieveRequest{}, nil)
|
||||
|
||||
if topK != 8 {
|
||||
t.Fatalf("expected fallback topK 8, got %d", topK)
|
||||
}
|
||||
if scoreThreshold != float32(0.3) {
|
||||
t.Fatalf("expected fallback score threshold 0.3, got %v", scoreThreshold)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package rag
|
||||
|
||||
type RetrieveRequest struct {
|
||||
KnowledgeBaseIDs []int64
|
||||
Query string
|
||||
TopK int
|
||||
ScoreThreshold float64
|
||||
}
|
||||
|
||||
type RetrieveResult struct {
|
||||
KnowledgeBaseID int64 `json:"knowledgeBaseId"`
|
||||
ChunkID int64 `json:"chunkId"`
|
||||
DocumentID int64 `json:"documentId"`
|
||||
DocumentTitle string `json:"documentTitle"`
|
||||
FaqID int64 `json:"faqId"`
|
||||
FaqQuestion string `json:"faqQuestion"`
|
||||
ChunkNo int `json:"chunkNo"`
|
||||
Title string `json:"title"`
|
||||
SectionPath string `json:"sectionPath"`
|
||||
Content string `json:"content"`
|
||||
Score float32 `json:"score"`
|
||||
ChunkType string `json:"chunkType"`
|
||||
}
|
||||
|
||||
type RerankRequest struct {
|
||||
Model string `json:"model"`
|
||||
Query string `json:"query"`
|
||||
Documents []string `json:"documents"`
|
||||
TopN int `json:"top_n"`
|
||||
}
|
||||
|
||||
type RerankResponse struct {
|
||||
Results []struct {
|
||||
Document string `json:"document"`
|
||||
Index int `json:"index"`
|
||||
RelevanceScore float64 `json:"relevance_score"`
|
||||
} `json:"results"`
|
||||
Meta struct {
|
||||
APIVersion struct {
|
||||
Version string `json:"version"`
|
||||
} `json:"api_version"`
|
||||
} `json:"meta"`
|
||||
}
|
||||
|
||||
type RerankResult struct {
|
||||
Index int `json:"index"`
|
||||
RelevanceScore float64 `json:"relevanceScore"`
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"strings"
|
||||
|
||||
"github.com/yuin/goldmark"
|
||||
"github.com/yuin/goldmark/extension"
|
||||
"github.com/yuin/goldmark/parser"
|
||||
"github.com/yuin/goldmark/renderer/html"
|
||||
htmlparser "golang.org/x/net/html"
|
||||
)
|
||||
|
||||
var plainTextMarkdown = goldmark.New(
|
||||
goldmark.WithExtensions(extension.GFM),
|
||||
goldmark.WithParserOptions(
|
||||
parser.WithAutoHeadingID(),
|
||||
),
|
||||
goldmark.WithRendererOptions(
|
||||
html.WithHardWraps(),
|
||||
html.WithXHTML(),
|
||||
),
|
||||
)
|
||||
|
||||
func ExtractPlainText(content string, contentType enums.KnowledgeDocumentContentType) string {
|
||||
switch contentType {
|
||||
case enums.KnowledgeDocumentContentTypeMarkdown:
|
||||
return ExtractPlainTextFromMarkdown(content)
|
||||
case enums.KnowledgeDocumentContentTypeHTML:
|
||||
return ExtractPlainTextFromHTML(content)
|
||||
default:
|
||||
return normalizeWhitespace(content)
|
||||
}
|
||||
}
|
||||
|
||||
func ExtractPlainTextFromMarkdown(content string) string {
|
||||
content = strings.TrimSpace(content)
|
||||
if content == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
var buf strings.Builder
|
||||
if err := plainTextMarkdown.Convert([]byte(content), &buf); err != nil {
|
||||
return normalizeWhitespace(content)
|
||||
}
|
||||
return ExtractPlainTextFromHTML(buf.String())
|
||||
}
|
||||
|
||||
func ExtractPlainTextFromHTML(content string) string {
|
||||
content = strings.TrimSpace(content)
|
||||
if content == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
var builder strings.Builder
|
||||
parent := &htmlparser.Node{
|
||||
Type: htmlparser.ElementNode,
|
||||
Data: "div",
|
||||
}
|
||||
nodes, err := htmlparser.ParseFragment(strings.NewReader(content), parent)
|
||||
if err == nil {
|
||||
for _, node := range nodes {
|
||||
writeHTMLNodeText(&builder, node)
|
||||
}
|
||||
return normalizeWhitespace(builder.String())
|
||||
}
|
||||
|
||||
// 兜底:部分输入在 ParseFragment 下会失败(例如不符合 fragment 规则或上下文不匹配)。
|
||||
// 这里用完整 HTML 解析保证可用性。
|
||||
doc, err := htmlparser.Parse(strings.NewReader("<div>" + content + "</div>"))
|
||||
if err != nil {
|
||||
return normalizeWhitespace(content)
|
||||
}
|
||||
writeHTMLNodeText(&builder, doc)
|
||||
return normalizeWhitespace(builder.String())
|
||||
}
|
||||
|
||||
func writeHTMLNodeText(builder *strings.Builder, node *htmlparser.Node) {
|
||||
if node == nil {
|
||||
return
|
||||
}
|
||||
|
||||
switch node.Type {
|
||||
case htmlparser.TextNode:
|
||||
builder.WriteString(node.Data)
|
||||
case htmlparser.ElementNode:
|
||||
if shouldSeparateHTMLText(node.Data) {
|
||||
builder.WriteByte(' ')
|
||||
}
|
||||
}
|
||||
|
||||
for child := node.FirstChild; child != nil; child = child.NextSibling {
|
||||
writeHTMLNodeText(builder, child)
|
||||
}
|
||||
|
||||
if node.Type == htmlparser.ElementNode && shouldSeparateHTMLText(node.Data) {
|
||||
builder.WriteByte(' ')
|
||||
}
|
||||
}
|
||||
|
||||
func shouldSeparateHTMLText(tag string) bool {
|
||||
switch tag {
|
||||
case "p", "div", "br", "li", "ul", "ol", "blockquote", "pre", "table", "tr", "td", "th", "h1", "h2", "h3", "h4", "h5", "h6":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeWhitespace(content string) string {
|
||||
return strings.Join(strings.Fields(strings.TrimSpace(content)), " ")
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package rag
|
||||
|
||||
import (
|
||||
"cs-agent/internal/ai/rag/vectordb"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestExtractPlainTextFromHTMLSeparatesBlockContent(t *testing.T) {
|
||||
got := ExtractPlainTextFromHTML("<div>Hello</div><div>World</div><p>Again</p>")
|
||||
want := "Hello World Again"
|
||||
if got != want {
|
||||
t.Fatalf("expected %q, got %q", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPlainTextMarkdownUsesGoldmark(t *testing.T) {
|
||||
got := ExtractPlainText("# Title\n\n- one\n- two", enums.KnowledgeDocumentContentTypeMarkdown)
|
||||
want := "Title one two"
|
||||
if got != want {
|
||||
t.Fatalf("expected %q, got %q", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChunkPayloadFromMapSupportsTypedConversion(t *testing.T) {
|
||||
got := vectordb.ChunkPayloadFromMap(map[string]any{
|
||||
"knowledge_base_id": "1",
|
||||
"document_id": "123",
|
||||
"document_title": "Doc",
|
||||
"chunk_no": "2",
|
||||
"chunk_type": "text",
|
||||
"section_path": "A > B",
|
||||
"title": "hello",
|
||||
"content": "world",
|
||||
"provider": "structured",
|
||||
})
|
||||
if got.KnowledgeBaseID != 1 || got.DocumentID != 123 || got.ChunkNo != 2 {
|
||||
t.Fatalf("unexpected numeric conversion result: %+v", got)
|
||||
}
|
||||
if got.DocumentTitle != "Doc" || got.SectionPath != "A > B" || got.Provider != "structured" {
|
||||
t.Fatalf("unexpected string conversion result: %+v", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package vectordb
|
||||
|
||||
import (
|
||||
"github.com/mlogclub/simple/common/structs"
|
||||
"github.com/spf13/cast"
|
||||
)
|
||||
|
||||
type ChunkPayload struct {
|
||||
KnowledgeBaseID int64 `json:"knowledge_base_id"`
|
||||
DocumentID int64 `json:"document_id"`
|
||||
DocumentTitle string `json:"document_title"`
|
||||
FaqID int64 `json:"faq_id"`
|
||||
FaqQuestion string `json:"faq_question"`
|
||||
ChunkNo int `json:"chunk_no"`
|
||||
ChunkType string `json:"chunk_type"`
|
||||
SectionPath string `json:"section_path"`
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
Provider string `json:"provider"`
|
||||
}
|
||||
|
||||
func (p ChunkPayload) ToMap() map[string]any {
|
||||
return structs.StructToMap(p)
|
||||
}
|
||||
|
||||
func ChunkPayloadFromMap(data map[string]any) ChunkPayload {
|
||||
if data == nil {
|
||||
return ChunkPayload{}
|
||||
}
|
||||
return ChunkPayload{
|
||||
KnowledgeBaseID: cast.ToInt64(data["knowledge_base_id"]),
|
||||
DocumentID: cast.ToInt64(data["document_id"]),
|
||||
DocumentTitle: cast.ToString(data["document_title"]),
|
||||
FaqID: cast.ToInt64(data["faq_id"]),
|
||||
FaqQuestion: cast.ToString(data["faq_question"]),
|
||||
ChunkNo: cast.ToInt(data["chunk_no"]),
|
||||
ChunkType: cast.ToString(data["chunk_type"]),
|
||||
SectionPath: cast.ToString(data["section_path"]),
|
||||
Title: cast.ToString(data["title"]),
|
||||
Content: cast.ToString(data["content"]),
|
||||
Provider: cast.ToString(data["provider"]),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package vectordb
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"cs-agent/internal/pkg/config"
|
||||
"cs-agent/internal/pkg/enums"
|
||||
)
|
||||
|
||||
var defaultProvider Provider
|
||||
|
||||
func Init(cfg *config.VectorDBConfig) error {
|
||||
if cfg == nil || cfg.Type == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
var err error
|
||||
switch enums.VectorDBType(cfg.Type) {
|
||||
case enums.VectorDBTypeQdrant:
|
||||
defaultProvider, err = NewQdrantProvider(cfg)
|
||||
default:
|
||||
return fmt.Errorf("unsupported vectordb type: %s", cfg.Type)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func GetProvider() Provider {
|
||||
return defaultProvider
|
||||
}
|
||||
|
||||
func Close() error {
|
||||
if defaultProvider != nil {
|
||||
return defaultProvider.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func CreateCollection(ctx context.Context, name string, dimension int) error {
|
||||
if defaultProvider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.CreateCollection(ctx, name, dimension)
|
||||
}
|
||||
|
||||
func DeleteCollection(ctx context.Context, name string) error {
|
||||
if defaultProvider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.DeleteCollection(ctx, name)
|
||||
}
|
||||
|
||||
func GetCollection(ctx context.Context, name string) (*CollectionInfo, error) {
|
||||
if defaultProvider == nil {
|
||||
return nil, fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.GetCollection(ctx, name)
|
||||
}
|
||||
|
||||
func ListCollections(ctx context.Context) ([]string, error) {
|
||||
if defaultProvider == nil {
|
||||
return nil, fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.ListCollections(ctx)
|
||||
}
|
||||
|
||||
func UpsertVectors(ctx context.Context, collectionName string, vectors []Vector) error {
|
||||
if defaultProvider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.UpsertVectors(ctx, collectionName, vectors)
|
||||
}
|
||||
|
||||
func DeleteVectors(ctx context.Context, collectionName string, ids []string) error {
|
||||
if defaultProvider == nil {
|
||||
return fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.DeleteVectors(ctx, collectionName, ids)
|
||||
}
|
||||
|
||||
func Search(ctx context.Context, req *SearchRequest) ([]SearchResult, error) {
|
||||
if defaultProvider == nil {
|
||||
return nil, fmt.Errorf("vectordb provider not initialized")
|
||||
}
|
||||
return defaultProvider.Search(ctx, req)
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
package vectordb
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/qdrant/go-client/qdrant"
|
||||
|
||||
"cs-agent/internal/pkg/config"
|
||||
)
|
||||
|
||||
type Vector struct {
|
||||
ID string `json:"id"`
|
||||
Vector []float32 `json:"vector"`
|
||||
Payload ChunkPayload `json:"payload"`
|
||||
}
|
||||
|
||||
type SearchRequest struct {
|
||||
CollectionName string `json:"collectionName"`
|
||||
Vector []float32 `json:"vector"`
|
||||
TopK int `json:"topK"`
|
||||
ScoreThreshold float32 `json:"scoreThreshold"`
|
||||
Filter *SearchFilter `json:"filter,omitempty"`
|
||||
}
|
||||
|
||||
type SearchFilter struct {
|
||||
KnowledgeBaseIDs []int64 `json:"knowledgeBaseIds,omitempty"`
|
||||
DocumentIDs []int64 `json:"documentIds,omitempty"`
|
||||
}
|
||||
|
||||
type SearchResult struct {
|
||||
ID string `json:"id"`
|
||||
Score float32 `json:"score"`
|
||||
Payload ChunkPayload `json:"payload"`
|
||||
}
|
||||
|
||||
type CollectionInfo struct {
|
||||
Name string `json:"name"`
|
||||
Dimension int `json:"dimension"`
|
||||
PointCount int `json:"pointCount"`
|
||||
Status string `json:"status"`
|
||||
}
|
||||
|
||||
type Provider interface {
|
||||
CreateCollection(ctx context.Context, name string, dimension int) error
|
||||
DeleteCollection(ctx context.Context, name string) error
|
||||
GetCollection(ctx context.Context, name string) (*CollectionInfo, error)
|
||||
ListCollections(ctx context.Context) ([]string, error)
|
||||
|
||||
UpsertVectors(ctx context.Context, collectionName string, vectors []Vector) error
|
||||
DeleteVectors(ctx context.Context, collectionName string, ids []string) error
|
||||
|
||||
Search(ctx context.Context, req *SearchRequest) ([]SearchResult, error)
|
||||
Close() error
|
||||
}
|
||||
|
||||
type QdrantProvider struct {
|
||||
client *qdrant.Client
|
||||
}
|
||||
|
||||
func NewQdrantProvider(cfg *config.VectorDBConfig) (*QdrantProvider, error) {
|
||||
if cfg == nil {
|
||||
return nil, fmt.Errorf("vectordb config is nil")
|
||||
}
|
||||
|
||||
host := cfg.Host
|
||||
if host == "" {
|
||||
host = "localhost"
|
||||
}
|
||||
|
||||
port := cfg.GrpcPort
|
||||
if port <= 0 {
|
||||
port = 6334
|
||||
}
|
||||
|
||||
client, err := qdrant.NewClient(&qdrant.Config{
|
||||
Host: host,
|
||||
Port: port,
|
||||
APIKey: cfg.APIKey,
|
||||
UseTLS: cfg.UseTLS,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create qdrant client: %w", err)
|
||||
}
|
||||
|
||||
return &QdrantProvider{client: client}, nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) Close() error {
|
||||
if p.client != nil {
|
||||
return p.client.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) CreateCollection(ctx context.Context, name string, dimension int) error {
|
||||
err := p.client.CreateCollection(ctx, &qdrant.CreateCollection{
|
||||
CollectionName: name,
|
||||
VectorsConfig: qdrant.NewVectorsConfig(&qdrant.VectorParams{
|
||||
Size: uint64(dimension),
|
||||
Distance: qdrant.Distance_Cosine,
|
||||
}),
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create collection %s: %w", name, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) DeleteCollection(ctx context.Context, name string) error {
|
||||
err := p.client.DeleteCollection(ctx, name)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete collection %s: %w", name, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) GetCollection(ctx context.Context, name string) (*CollectionInfo, error) {
|
||||
info, err := p.client.GetCollectionInfo(ctx, name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get collection %s: %w", name, err)
|
||||
}
|
||||
|
||||
status := info.GetStatus().String()
|
||||
pointCount := int(info.GetPointsCount())
|
||||
|
||||
dimension := 0
|
||||
if info.Config != nil && info.Config.Params != nil {
|
||||
vectorsConfig := info.Config.Params.VectorsConfig
|
||||
if vectorsConfig != nil {
|
||||
params := vectorsConfig.GetParams()
|
||||
if params != nil {
|
||||
dimension = int(params.Size)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return &CollectionInfo{
|
||||
Name: name,
|
||||
Dimension: dimension,
|
||||
PointCount: pointCount,
|
||||
Status: status,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) ListCollections(ctx context.Context) ([]string, error) {
|
||||
collections, err := p.client.ListCollections(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to list collections: %w", err)
|
||||
}
|
||||
|
||||
return collections, nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) UpsertVectors(ctx context.Context, collectionName string, vectors []Vector) error {
|
||||
if len(vectors) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
points := make([]*qdrant.PointStruct, 0, len(vectors))
|
||||
for _, v := range vectors {
|
||||
points = append(points, &qdrant.PointStruct{
|
||||
Id: qdrant.NewID(v.ID),
|
||||
Vectors: qdrant.NewVectors(v.Vector...),
|
||||
Payload: qdrant.NewValueMap(v.Payload.ToMap()),
|
||||
})
|
||||
}
|
||||
|
||||
_, err := p.client.Upsert(ctx, &qdrant.UpsertPoints{
|
||||
CollectionName: collectionName,
|
||||
Points: points,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to upsert vectors to collection %s: %w", collectionName, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) DeleteVectors(ctx context.Context, collectionName string, ids []string) error {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
pointIDs := make([]*qdrant.PointId, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
pointIDs = append(pointIDs, qdrant.NewID(id))
|
||||
}
|
||||
|
||||
_, err := p.client.Delete(ctx, &qdrant.DeletePoints{
|
||||
CollectionName: collectionName,
|
||||
Points: &qdrant.PointsSelector{
|
||||
PointsSelectorOneOf: &qdrant.PointsSelector_Points{
|
||||
Points: &qdrant.PointsIdsList{
|
||||
Ids: pointIDs,
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete vectors from collection %s: %w", collectionName, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) Search(ctx context.Context, req *SearchRequest) ([]SearchResult, error) {
|
||||
filter := p.buildFilter(req.Filter)
|
||||
|
||||
results, err := p.client.Query(ctx, &qdrant.QueryPoints{
|
||||
CollectionName: req.CollectionName,
|
||||
Query: qdrant.NewQuery(req.Vector...),
|
||||
Limit: qdrant.PtrOf(uint64(req.TopK)),
|
||||
ScoreThreshold: &req.ScoreThreshold,
|
||||
Filter: filter,
|
||||
WithPayload: qdrant.NewWithPayload(true),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to search collection %s: %w", req.CollectionName, err)
|
||||
}
|
||||
|
||||
searchResults := make([]SearchResult, 0, len(results))
|
||||
for _, r := range results {
|
||||
payload := make(map[string]any)
|
||||
if r.Payload != nil {
|
||||
for k, v := range r.Payload {
|
||||
payload[k] = p.extractPayloadValue(v)
|
||||
}
|
||||
}
|
||||
|
||||
id := ""
|
||||
if r.Id != nil {
|
||||
id = r.Id.GetUuid()
|
||||
}
|
||||
|
||||
searchResults = append(searchResults, SearchResult{
|
||||
ID: id,
|
||||
Score: r.Score,
|
||||
Payload: ChunkPayloadFromMap(payload),
|
||||
})
|
||||
}
|
||||
|
||||
return searchResults, nil
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) buildFilter(filter *SearchFilter) *qdrant.Filter {
|
||||
if filter == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
must := make([]*qdrant.Condition, 0, 2)
|
||||
if len(filter.KnowledgeBaseIDs) > 0 {
|
||||
must = append(must, qdrant.NewMatchInts("knowledge_base_id", filter.KnowledgeBaseIDs...))
|
||||
}
|
||||
if len(filter.DocumentIDs) > 0 {
|
||||
must = append(must, qdrant.NewMatchInts("document_id", filter.DocumentIDs...))
|
||||
}
|
||||
if len(must) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &qdrant.Filter{Must: must}
|
||||
}
|
||||
|
||||
func (p *QdrantProvider) extractPayloadValue(v *qdrant.Value) interface{} {
|
||||
if v == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch val := v.Kind.(type) {
|
||||
case *qdrant.Value_StringValue:
|
||||
return val.StringValue
|
||||
case *qdrant.Value_IntegerValue:
|
||||
return val.IntegerValue
|
||||
case *qdrant.Value_DoubleValue:
|
||||
return val.DoubleValue
|
||||
case *qdrant.Value_BoolValue:
|
||||
return val.BoolValue
|
||||
case *qdrant.Value_ListValue:
|
||||
list := make([]interface{}, 0, len(val.ListValue.Values))
|
||||
for _, item := range val.ListValue.Values {
|
||||
list = append(list, p.extractPayloadValue(item))
|
||||
}
|
||||
return list
|
||||
case *qdrant.Value_StructValue:
|
||||
m := make(map[string]interface{})
|
||||
for k, v := range val.StructValue.Fields {
|
||||
m[k] = p.extractPayloadValue(v)
|
||||
}
|
||||
return m
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user