This commit is contained in:
mlogclub
2026-04-09 10:01:23 +08:00
commit efe801b8bf
707 changed files with 110595 additions and 0 deletions
+432
View File
@@ -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 "未知"
}
}
+128
View File
@@ -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)
}
}
+44
View File
@@ -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
}
+12
View File
@@ -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)
}
+59
View File
@@ -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, " ")
}
+32
View File
@@ -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
}
+218
View File
@@ -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
}
+702
View File
@@ -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
}
+140
View File
@@ -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
}
+507
View File
@@ -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"`
}
+300
View File
@@ -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
}
+49
View File
@@ -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)
}
}
+48
View File
@@ -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"`
}
+112
View File
@@ -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)), " ")
}
+43
View File
@@ -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)
}
}
+43
View File
@@ -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"]),
}
}
+86
View File
@@ -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)
}
+292
View File
@@ -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
}
}