feat: enhance knowledge retrieval by adding MaxContextItems support and refactoring related methods
This commit is contained in:
@@ -159,9 +159,9 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) {
|
|||||||
messages = append(messages, history.Messages...)
|
messages = append(messages, history.Messages...)
|
||||||
|
|
||||||
retriever := retrievers.NewKnowledgeRetriever(req.AIAgent)
|
retriever := retrievers.NewKnowledgeRetriever(req.AIAgent)
|
||||||
if retrieveResult, retrieveErr := retriever.RetrieveContextByOptions(ctx, retrievers.KnowledgeRetrieveOptions{
|
retrieveOptions := retrievers.DefaultKnowledgeRetrieveOptions()
|
||||||
QueryPreview: preview(req.UserMessage.Content, 120),
|
retrieveOptions.QueryPreview = preview(req.UserMessage.Content, 120)
|
||||||
}, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil && retrieveResult != nil {
|
if retrieveResult, retrieveErr := retriever.RetrieveContextByOptions(ctx, retrieveOptions, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil && retrieveResult != nil {
|
||||||
summary.RetrieverCount = len(retrieveResult.Hits)
|
summary.RetrieverCount = len(retrieveResult.Hits)
|
||||||
collector.SetRetrieverSummary(retrieveResult.TraceSummary)
|
collector.SetRetrieverSummary(retrieveResult.TraceSummary)
|
||||||
collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, retrieveResult.TraceItems...)
|
collector.Data.Retriever.Items = append(collector.Data.Retriever.Items, retrieveResult.TraceItems...)
|
||||||
|
|||||||
@@ -74,6 +74,7 @@ func (c *RuntimeTraceCollector) SetRetrieverSummary(summary RetrieverTraceSummar
|
|||||||
c.Data.Retriever.TopK = summary.TopK
|
c.Data.Retriever.TopK = summary.TopK
|
||||||
c.Data.Retriever.ScoreThreshold = summary.ScoreThreshold
|
c.Data.Retriever.ScoreThreshold = summary.ScoreThreshold
|
||||||
c.Data.Retriever.ContextMaxTokens = summary.ContextMaxTokens
|
c.Data.Retriever.ContextMaxTokens = summary.ContextMaxTokens
|
||||||
|
c.Data.Retriever.MaxContextItems = summary.MaxContextItems
|
||||||
c.Data.Retriever.Count = summary.HitCount
|
c.Data.Retriever.Count = summary.HitCount
|
||||||
c.Data.Retriever.ContextCount = summary.ContextCount
|
c.Data.Retriever.ContextCount = summary.ContextCount
|
||||||
c.Data.Retriever.EmbeddingMs = summary.EmbeddingMs
|
c.Data.Retriever.EmbeddingMs = summary.EmbeddingMs
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ type RetrieverTraceSummary struct {
|
|||||||
TopK int
|
TopK int
|
||||||
ScoreThreshold float64
|
ScoreThreshold float64
|
||||||
ContextMaxTokens int
|
ContextMaxTokens int
|
||||||
|
MaxContextItems int
|
||||||
HitCount int
|
HitCount int
|
||||||
ContextCount int
|
ContextCount int
|
||||||
EmbeddingMs int64
|
EmbeddingMs int64
|
||||||
@@ -103,6 +104,7 @@ type RuntimeTraceData struct {
|
|||||||
TopK int `json:"topK,omitempty"`
|
TopK int `json:"topK,omitempty"`
|
||||||
ScoreThreshold float64 `json:"scoreThreshold,omitempty"`
|
ScoreThreshold float64 `json:"scoreThreshold,omitempty"`
|
||||||
ContextMaxTokens int `json:"contextMaxTokens,omitempty"`
|
ContextMaxTokens int `json:"contextMaxTokens,omitempty"`
|
||||||
|
MaxContextItems int `json:"maxContextItems,omitempty"`
|
||||||
ContextCount int `json:"contextCount,omitempty"`
|
ContextCount int `json:"contextCount,omitempty"`
|
||||||
EmbeddingMs int64 `json:"embeddingMs,omitempty"`
|
EmbeddingMs int64 `json:"embeddingMs,omitempty"`
|
||||||
VectorSearchMs int64 `json:"vectorSearchMs,omitempty"`
|
VectorSearchMs int64 `json:"vectorSearchMs,omitempty"`
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ import (
|
|||||||
const defaultRuntimeKnowledgeContextTokens = 4000
|
const defaultRuntimeKnowledgeContextTokens = 4000
|
||||||
const defaultRuntimeKnowledgeTopK = 8
|
const defaultRuntimeKnowledgeTopK = 8
|
||||||
const defaultRuntimeKnowledgeScoreThreshold = 0.3
|
const defaultRuntimeKnowledgeScoreThreshold = 0.3
|
||||||
|
const defaultRuntimeKnowledgeMaxContextItems = 5
|
||||||
|
|
||||||
type KnowledgeRetriever struct {
|
type KnowledgeRetriever struct {
|
||||||
AIAgent *models.AIAgent
|
AIAgent *models.AIAgent
|
||||||
@@ -24,6 +25,7 @@ type KnowledgeRetriever struct {
|
|||||||
|
|
||||||
type KnowledgeRetrieveOptions struct {
|
type KnowledgeRetrieveOptions struct {
|
||||||
ContextMaxTokens int
|
ContextMaxTokens int
|
||||||
|
MaxContextItems int
|
||||||
TopK int
|
TopK int
|
||||||
ScoreThreshold float64
|
ScoreThreshold float64
|
||||||
QueryPreview string
|
QueryPreview string
|
||||||
@@ -52,6 +54,13 @@ func NewKnowledgeRetriever(aiAgent *models.AIAgent) *KnowledgeRetriever {
|
|||||||
return &KnowledgeRetriever{AIAgent: aiAgent}
|
return &KnowledgeRetriever{AIAgent: aiAgent}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func DefaultKnowledgeRetrieveOptions() KnowledgeRetrieveOptions {
|
||||||
|
return KnowledgeRetrieveOptions{
|
||||||
|
ContextMaxTokens: defaultRuntimeKnowledgeContextTokens,
|
||||||
|
MaxContextItems: defaultRuntimeKnowledgeMaxContextItems,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (r *KnowledgeRetriever) KnowledgeBaseIDs() []int64 {
|
func (r *KnowledgeRetriever) KnowledgeBaseIDs() []int64 {
|
||||||
if r == nil || r.AIAgent == nil {
|
if r == nil || r.AIAgent == nil {
|
||||||
return nil
|
return nil
|
||||||
@@ -60,7 +69,7 @@ func (r *KnowledgeRetriever) KnowledgeBaseIDs() []int64 {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) {
|
func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) {
|
||||||
return r.RetrieveByOptions(ctx, KnowledgeRetrieveOptions{}, query)
|
return r.RetrieveByOptions(ctx, DefaultKnowledgeRetrieveOptions(), query)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *KnowledgeRetriever) RetrieveByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) {
|
func (r *KnowledgeRetriever) RetrieveByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) {
|
||||||
@@ -74,7 +83,7 @@ func (r *KnowledgeRetriever) RetrieveByOptions(ctx context.Context, opts Knowled
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string) (*KnowledgeRetrieveResult, error) {
|
func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string) (*KnowledgeRetrieveResult, error) {
|
||||||
return r.RetrieveContextByOptions(ctx, KnowledgeRetrieveOptions{}, query)
|
return r.RetrieveContextByOptions(ctx, DefaultKnowledgeRetrieveOptions(), query)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) (*KnowledgeRetrieveResult, error) {
|
func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) (*KnowledgeRetrieveResult, error) {
|
||||||
@@ -85,6 +94,10 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts
|
|||||||
if contextMaxTokens <= 0 {
|
if contextMaxTokens <= 0 {
|
||||||
contextMaxTokens = defaultRuntimeKnowledgeContextTokens
|
contextMaxTokens = defaultRuntimeKnowledgeContextTokens
|
||||||
}
|
}
|
||||||
|
maxContextItems := opts.MaxContextItems
|
||||||
|
if maxContextItems <= 0 {
|
||||||
|
maxContextItems = defaultRuntimeKnowledgeMaxContextItems
|
||||||
|
}
|
||||||
queryPreview := strings.TrimSpace(opts.QueryPreview)
|
queryPreview := strings.TrimSpace(opts.QueryPreview)
|
||||||
if queryPreview == "" {
|
if queryPreview == "" {
|
||||||
queryPreview = query
|
queryPreview = query
|
||||||
@@ -94,6 +107,7 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts
|
|||||||
Query: query,
|
Query: query,
|
||||||
Options: KnowledgeRetrieveOptions{
|
Options: KnowledgeRetrieveOptions{
|
||||||
ContextMaxTokens: contextMaxTokens,
|
ContextMaxTokens: contextMaxTokens,
|
||||||
|
MaxContextItems: maxContextItems,
|
||||||
TopK: opts.TopK,
|
TopK: opts.TopK,
|
||||||
ScoreThreshold: opts.ScoreThreshold,
|
ScoreThreshold: opts.ScoreThreshold,
|
||||||
QueryPreview: queryPreview,
|
QueryPreview: queryPreview,
|
||||||
@@ -110,12 +124,30 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts
|
|||||||
ret.Hits = append([]rag.RetrieveResult(nil), results...)
|
ret.Hits = append([]rag.RetrieveResult(nil), results...)
|
||||||
ret.Trace = trace
|
ret.Trace = trace
|
||||||
ret.ContextResults = rag.Retrieve.SelectContextResults(results, contextMaxTokens)
|
ret.ContextResults = rag.Retrieve.SelectContextResults(results, contextMaxTokens)
|
||||||
ret.ContextText = strings.TrimSpace(rag.Retrieve.BuildContext(ctx, results, contextMaxTokens))
|
ret.ContextResults = limitContextResults(ret.ContextResults, maxContextItems)
|
||||||
|
ret.ContextText = strings.TrimSpace(buildContextText(ret.ContextResults))
|
||||||
ret.TraceItems = buildRetrieverTraceItems(queryPreview, results, trace)
|
ret.TraceItems = buildRetrieverTraceItems(queryPreview, results, trace)
|
||||||
ret.TraceSummary = buildRetrieverTraceSummary(ret.Options, ret.Policies, ret.ContextResults, results, trace)
|
ret.TraceSummary = buildRetrieverTraceSummary(ret.Options, ret.Policies, ret.ContextResults, results, trace)
|
||||||
return ret, nil
|
return ret, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func limitContextResults(results []rag.RetrieveResult, maxItems int) []rag.RetrieveResult {
|
||||||
|
if len(results) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if maxItems <= 0 || len(results) <= maxItems {
|
||||||
|
return append([]rag.RetrieveResult(nil), results...)
|
||||||
|
}
|
||||||
|
return append([]rag.RetrieveResult(nil), results[:maxItems]...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildContextText(results []rag.RetrieveResult) string {
|
||||||
|
if len(results) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(rag.Retrieve.BuildContext(context.Background(), results, 1<<30))
|
||||||
|
}
|
||||||
|
|
||||||
func (r *KnowledgeRetriever) resolvePolicies(knowledgeBaseIDs []int64, opts KnowledgeRetrieveOptions) []KnowledgeBaseRetrievePolicy {
|
func (r *KnowledgeRetriever) resolvePolicies(knowledgeBaseIDs []int64, opts KnowledgeRetrieveOptions) []KnowledgeBaseRetrievePolicy {
|
||||||
if len(knowledgeBaseIDs) == 0 {
|
if len(knowledgeBaseIDs) == 0 {
|
||||||
return nil
|
return nil
|
||||||
@@ -192,6 +224,7 @@ func buildRetrieverTraceSummary(opts KnowledgeRetrieveOptions, policies []Knowle
|
|||||||
TopK: opts.TopK,
|
TopK: opts.TopK,
|
||||||
ScoreThreshold: opts.ScoreThreshold,
|
ScoreThreshold: opts.ScoreThreshold,
|
||||||
ContextMaxTokens: opts.ContextMaxTokens,
|
ContextMaxTokens: opts.ContextMaxTokens,
|
||||||
|
MaxContextItems: opts.MaxContextItems,
|
||||||
HitCount: len(results),
|
HitCount: len(results),
|
||||||
ContextCount: len(contextResults),
|
ContextCount: len(contextResults),
|
||||||
Policies: buildRetrieverPolicyTraceItems(policies),
|
Policies: buildRetrieverPolicyTraceItems(policies),
|
||||||
|
|||||||
Reference in New Issue
Block a user