From a5cf31656eab211d06adad4d4d12e1ba4a22b3a8 Mon Sep 17 00:00:00 2001 From: mlogclub Date: Mon, 13 Apr 2026 20:01:38 +0800 Subject: [PATCH] feat: implement context handling methods for improved result selection and context building --- internal/ai/rag/retrieve.go | 144 -------------------------- internal/ai/rag/retrieve_context.go | 150 ++++++++++++++++++++++++++++ 2 files changed, 150 insertions(+), 144 deletions(-) create mode 100644 internal/ai/rag/retrieve_context.go diff --git a/internal/ai/rag/retrieve.go b/internal/ai/rag/retrieve.go index aa6298a..9dc0191 100644 --- a/internal/ai/rag/retrieve.go +++ b/internal/ai/rag/retrieve.go @@ -154,150 +154,6 @@ func (s *retrieve) rerank(ctx context.Context, query string, results []RetrieveR 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 { diff --git a/internal/ai/rag/retrieve_context.go b/internal/ai/rag/retrieve_context.go new file mode 100644 index 0000000..7aaf61e --- /dev/null +++ b/internal/ai/rag/retrieve_context.go @@ -0,0 +1,150 @@ +package rag + +import ( + "context" + "fmt" + "strings" +) + +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(_ context.Context, results []RetrieveResult, maxTokens int) string { + if len(results) == 0 { + return "" + } + + normalizedResults := s.SelectContextResults(results, maxTokens) + var builder strings.Builder + for _, r := range normalizedResults { + builder.WriteString(buildContextChunkText(r)) + } + + return builder.String() +} + +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) +}