feat: enhance knowledge retrieval by adding MaxContextItems support and refactoring related methods
This commit is contained in:
@@ -74,6 +74,7 @@ func (c *RuntimeTraceCollector) SetRetrieverSummary(summary RetrieverTraceSummar
|
||||
c.Data.Retriever.TopK = summary.TopK
|
||||
c.Data.Retriever.ScoreThreshold = summary.ScoreThreshold
|
||||
c.Data.Retriever.ContextMaxTokens = summary.ContextMaxTokens
|
||||
c.Data.Retriever.MaxContextItems = summary.MaxContextItems
|
||||
c.Data.Retriever.Count = summary.HitCount
|
||||
c.Data.Retriever.ContextCount = summary.ContextCount
|
||||
c.Data.Retriever.EmbeddingMs = summary.EmbeddingMs
|
||||
|
||||
@@ -45,6 +45,7 @@ type RetrieverTraceSummary struct {
|
||||
TopK int
|
||||
ScoreThreshold float64
|
||||
ContextMaxTokens int
|
||||
MaxContextItems int
|
||||
HitCount int
|
||||
ContextCount int
|
||||
EmbeddingMs int64
|
||||
@@ -103,6 +104,7 @@ type RuntimeTraceData struct {
|
||||
TopK int `json:"topK,omitempty"`
|
||||
ScoreThreshold float64 `json:"scoreThreshold,omitempty"`
|
||||
ContextMaxTokens int `json:"contextMaxTokens,omitempty"`
|
||||
MaxContextItems int `json:"maxContextItems,omitempty"`
|
||||
ContextCount int `json:"contextCount,omitempty"`
|
||||
EmbeddingMs int64 `json:"embeddingMs,omitempty"`
|
||||
VectorSearchMs int64 `json:"vectorSearchMs,omitempty"`
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
const defaultRuntimeKnowledgeContextTokens = 4000
|
||||
const defaultRuntimeKnowledgeTopK = 8
|
||||
const defaultRuntimeKnowledgeScoreThreshold = 0.3
|
||||
const defaultRuntimeKnowledgeMaxContextItems = 5
|
||||
|
||||
type KnowledgeRetriever struct {
|
||||
AIAgent *models.AIAgent
|
||||
@@ -24,6 +25,7 @@ type KnowledgeRetriever struct {
|
||||
|
||||
type KnowledgeRetrieveOptions struct {
|
||||
ContextMaxTokens int
|
||||
MaxContextItems int
|
||||
TopK int
|
||||
ScoreThreshold float64
|
||||
QueryPreview string
|
||||
@@ -52,6 +54,13 @@ func NewKnowledgeRetriever(aiAgent *models.AIAgent) *KnowledgeRetriever {
|
||||
return &KnowledgeRetriever{AIAgent: aiAgent}
|
||||
}
|
||||
|
||||
func DefaultKnowledgeRetrieveOptions() KnowledgeRetrieveOptions {
|
||||
return KnowledgeRetrieveOptions{
|
||||
ContextMaxTokens: defaultRuntimeKnowledgeContextTokens,
|
||||
MaxContextItems: defaultRuntimeKnowledgeMaxContextItems,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *KnowledgeRetriever) KnowledgeBaseIDs() []int64 {
|
||||
if r == nil || r.AIAgent == 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) {
|
||||
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) {
|
||||
@@ -74,7 +83,7 @@ func (r *KnowledgeRetriever) RetrieveByOptions(ctx context.Context, opts Knowled
|
||||
}
|
||||
|
||||
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) {
|
||||
@@ -85,6 +94,10 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts
|
||||
if contextMaxTokens <= 0 {
|
||||
contextMaxTokens = defaultRuntimeKnowledgeContextTokens
|
||||
}
|
||||
maxContextItems := opts.MaxContextItems
|
||||
if maxContextItems <= 0 {
|
||||
maxContextItems = defaultRuntimeKnowledgeMaxContextItems
|
||||
}
|
||||
queryPreview := strings.TrimSpace(opts.QueryPreview)
|
||||
if queryPreview == "" {
|
||||
queryPreview = query
|
||||
@@ -94,6 +107,7 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts
|
||||
Query: query,
|
||||
Options: KnowledgeRetrieveOptions{
|
||||
ContextMaxTokens: contextMaxTokens,
|
||||
MaxContextItems: maxContextItems,
|
||||
TopK: opts.TopK,
|
||||
ScoreThreshold: opts.ScoreThreshold,
|
||||
QueryPreview: queryPreview,
|
||||
@@ -110,12 +124,30 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts
|
||||
ret.Hits = append([]rag.RetrieveResult(nil), results...)
|
||||
ret.Trace = trace
|
||||
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.TraceSummary = buildRetrieverTraceSummary(ret.Options, ret.Policies, ret.ContextResults, results, trace)
|
||||
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 {
|
||||
if len(knowledgeBaseIDs) == 0 {
|
||||
return nil
|
||||
@@ -192,6 +224,7 @@ func buildRetrieverTraceSummary(opts KnowledgeRetrieveOptions, policies []Knowle
|
||||
TopK: opts.TopK,
|
||||
ScoreThreshold: opts.ScoreThreshold,
|
||||
ContextMaxTokens: opts.ContextMaxTokens,
|
||||
MaxContextItems: opts.MaxContextItems,
|
||||
HitCount: len(results),
|
||||
ContextCount: len(contextResults),
|
||||
Policies: buildRetrieverPolicyTraceItems(policies),
|
||||
|
||||
Reference in New Issue
Block a user