feat: enhance knowledge retrieval by adding MaxContextItems support and refactoring related methods

This commit is contained in:
mlogclub
2026-04-12 11:02:08 +08:00
parent 0b77d91dfe
commit 0db5f0d663
4 changed files with 42 additions and 6 deletions
@@ -159,9 +159,9 @@ func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) {
messages = append(messages, history.Messages...)
retriever := retrievers.NewKnowledgeRetriever(req.AIAgent)
if retrieveResult, retrieveErr := retriever.RetrieveContextByOptions(ctx, retrievers.KnowledgeRetrieveOptions{
QueryPreview: preview(req.UserMessage.Content, 120),
}, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil && retrieveResult != nil {
retrieveOptions := retrievers.DefaultKnowledgeRetrieveOptions()
retrieveOptions.QueryPreview = preview(req.UserMessage.Content, 120)
if retrieveResult, retrieveErr := retriever.RetrieveContextByOptions(ctx, retrieveOptions, strings.TrimSpace(req.UserMessage.Content)); retrieveErr == nil && retrieveResult != nil {
summary.RetrieverCount = len(retrieveResult.Hits)
collector.SetRetrieverSummary(retrieveResult.TraceSummary)
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.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),