diff --git a/internal/ai/runtime/internal/engine/service.go b/internal/ai/runtime/internal/engine/service.go index 1a1d825..ae3e016 100644 --- a/internal/ai/runtime/internal/engine/service.go +++ b/internal/ai/runtime/internal/engine/service.go @@ -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...) diff --git a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go index 4e00604..4e39e1f 100644 --- a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go +++ b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go @@ -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 diff --git a/internal/ai/runtime/internal/impl/callbacks/trace_callback.go b/internal/ai/runtime/internal/impl/callbacks/trace_callback.go index 9410072..bb5d686 100644 --- a/internal/ai/runtime/internal/impl/callbacks/trace_callback.go +++ b/internal/ai/runtime/internal/impl/callbacks/trace_callback.go @@ -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"` diff --git a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go index 6e50e75..eca5317 100644 --- a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go +++ b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go @@ -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),