diff --git a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go index e1211dd..4e00604 100644 --- a/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go +++ b/internal/ai/runtime/internal/impl/callbacks/runlog_callback.go @@ -79,6 +79,7 @@ func (c *RuntimeTraceCollector) SetRetrieverSummary(summary RetrieverTraceSummar c.Data.Retriever.EmbeddingMs = summary.EmbeddingMs c.Data.Retriever.VectorSearchMs = summary.VectorSearchMs c.Data.Retriever.HydrateMs = summary.HydrateMs + c.Data.Retriever.Policies = append([]RetrieverPolicyTraceItem(nil), summary.Policies...) } func (c *RuntimeTraceCollector) AddToolItem(item ToolTraceItem) { diff --git a/internal/ai/runtime/internal/impl/callbacks/trace_callback.go b/internal/ai/runtime/internal/impl/callbacks/trace_callback.go index 5e45c3f..9410072 100644 --- a/internal/ai/runtime/internal/impl/callbacks/trace_callback.go +++ b/internal/ai/runtime/internal/impl/callbacks/trace_callback.go @@ -50,6 +50,13 @@ type RetrieverTraceSummary struct { EmbeddingMs int64 VectorSearchMs int64 HydrateMs int64 + Policies []RetrieverPolicyTraceItem +} + +type RetrieverPolicyTraceItem struct { + KnowledgeBaseID int64 `json:"knowledgeBaseId,omitempty"` + TopK int `json:"topK,omitempty"` + ScoreThreshold float64 `json:"scoreThreshold,omitempty"` } type InstructionTraceSummary struct { @@ -92,15 +99,16 @@ type RuntimeTraceData struct { CurrentUserMessagePreview string `json:"currentUserMessagePreview,omitempty"` } `json:"input"` Retriever struct { - Count int `json:"count,omitempty"` - TopK int `json:"topK,omitempty"` - ScoreThreshold float64 `json:"scoreThreshold,omitempty"` - ContextMaxTokens int `json:"contextMaxTokens,omitempty"` - ContextCount int `json:"contextCount,omitempty"` - EmbeddingMs int64 `json:"embeddingMs,omitempty"` - VectorSearchMs int64 `json:"vectorSearchMs,omitempty"` - HydrateMs int64 `json:"hydrateMs,omitempty"` - Items []RetrieverTraceItem `json:"items,omitempty"` + Count int `json:"count,omitempty"` + TopK int `json:"topK,omitempty"` + ScoreThreshold float64 `json:"scoreThreshold,omitempty"` + ContextMaxTokens int `json:"contextMaxTokens,omitempty"` + ContextCount int `json:"contextCount,omitempty"` + EmbeddingMs int64 `json:"embeddingMs,omitempty"` + VectorSearchMs int64 `json:"vectorSearchMs,omitempty"` + HydrateMs int64 `json:"hydrateMs,omitempty"` + Policies []RetrieverPolicyTraceItem `json:"policies,omitempty"` + Items []RetrieverTraceItem `json:"items,omitempty"` } `json:"retriever"` Tools struct { Count int `json:"count,omitempty"` diff --git a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go index 901c70a..6e50e75 100644 --- a/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go +++ b/internal/ai/runtime/internal/impl/retrievers/knowledge_retriever.go @@ -7,10 +7,16 @@ import ( "cs-agent/internal/ai/rag" "cs-agent/internal/ai/runtime/internal/impl/callbacks" "cs-agent/internal/models" + "cs-agent/internal/pkg/enums" "cs-agent/internal/pkg/utils" + "cs-agent/internal/repositories" + + "github.com/mlogclub/simple/sqls" ) const defaultRuntimeKnowledgeContextTokens = 4000 +const defaultRuntimeKnowledgeTopK = 8 +const defaultRuntimeKnowledgeScoreThreshold = 0.3 type KnowledgeRetriever struct { AIAgent *models.AIAgent @@ -23,6 +29,12 @@ type KnowledgeRetrieveOptions struct { QueryPreview string } +type KnowledgeBaseRetrievePolicy struct { + KnowledgeBaseID int64 + TopK int + ScoreThreshold float64 +} + type KnowledgeRetrieveResult struct { KnowledgeBaseIDs []int64 Query string @@ -33,6 +45,7 @@ type KnowledgeRetrieveResult struct { Trace *rag.RetrieveTrace TraceItems []callbacks.RetrieverTraceItem TraceSummary callbacks.RetrieverTraceSummary + Policies []KnowledgeBaseRetrievePolicy } func NewKnowledgeRetriever(aiAgent *models.AIAgent) *KnowledgeRetriever { @@ -67,6 +80,7 @@ func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string) func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) (*KnowledgeRetrieveResult, error) { query = strings.TrimSpace(query) knowledgeBaseIDs := r.KnowledgeBaseIDs() + policies := r.resolvePolicies(knowledgeBaseIDs, opts) contextMaxTokens := opts.ContextMaxTokens if contextMaxTokens <= 0 { contextMaxTokens = defaultRuntimeKnowledgeContextTokens @@ -84,6 +98,7 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts ScoreThreshold: opts.ScoreThreshold, QueryPreview: queryPreview, }, + Policies: append([]KnowledgeBaseRetrievePolicy(nil), policies...), } if query == "" || len(knowledgeBaseIDs) == 0 { return ret, nil @@ -97,10 +112,59 @@ func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts ret.ContextResults = rag.Retrieve.SelectContextResults(results, contextMaxTokens) ret.ContextText = strings.TrimSpace(rag.Retrieve.BuildContext(ctx, results, contextMaxTokens)) ret.TraceItems = buildRetrieverTraceItems(queryPreview, results, trace) - ret.TraceSummary = buildRetrieverTraceSummary(ret.Options, ret.ContextResults, results, trace) + ret.TraceSummary = buildRetrieverTraceSummary(ret.Options, ret.Policies, ret.ContextResults, results, trace) return ret, nil } +func (r *KnowledgeRetriever) resolvePolicies(knowledgeBaseIDs []int64, opts KnowledgeRetrieveOptions) []KnowledgeBaseRetrievePolicy { + if len(knowledgeBaseIDs) == 0 { + return nil + } + knowledgeBases := loadRuntimeKnowledgeBases(knowledgeBaseIDs) + ret := make([]KnowledgeBaseRetrievePolicy, 0, len(knowledgeBaseIDs)) + for _, knowledgeBaseID := range knowledgeBaseIDs { + policy := KnowledgeBaseRetrievePolicy{ + KnowledgeBaseID: knowledgeBaseID, + TopK: defaultRuntimeKnowledgeTopK, + ScoreThreshold: defaultRuntimeKnowledgeScoreThreshold, + } + if knowledgeBase, ok := knowledgeBases[knowledgeBaseID]; ok { + if knowledgeBase.DefaultTopK > 0 { + policy.TopK = knowledgeBase.DefaultTopK + } + if knowledgeBase.DefaultScoreThreshold > 0 { + policy.ScoreThreshold = knowledgeBase.DefaultScoreThreshold + } + } + if opts.TopK > 0 { + policy.TopK = opts.TopK + } + if opts.ScoreThreshold > 0 { + policy.ScoreThreshold = opts.ScoreThreshold + } + ret = append(ret, policy) + } + return ret +} + +func loadRuntimeKnowledgeBases(ids []int64) map[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 + } + ret := make(map[int64]models.KnowledgeBase, len(items)) + for _, item := range items { + if item.Status != enums.StatusOk { + continue + } + ret[item.ID] = item + } + return ret +} + func buildRetrieverTraceItems(queryPreview string, results []rag.RetrieveResult, trace *rag.RetrieveTrace) []callbacks.RetrieverTraceItem { if len(results) == 0 { return nil @@ -123,13 +187,20 @@ func buildRetrieverTraceItems(queryPreview string, results []rag.RetrieveResult, return ret } -func buildRetrieverTraceSummary(opts KnowledgeRetrieveOptions, contextResults []rag.RetrieveResult, results []rag.RetrieveResult, trace *rag.RetrieveTrace) callbacks.RetrieverTraceSummary { +func buildRetrieverTraceSummary(opts KnowledgeRetrieveOptions, policies []KnowledgeBaseRetrievePolicy, contextResults []rag.RetrieveResult, results []rag.RetrieveResult, trace *rag.RetrieveTrace) callbacks.RetrieverTraceSummary { ret := callbacks.RetrieverTraceSummary{ TopK: opts.TopK, ScoreThreshold: opts.ScoreThreshold, ContextMaxTokens: opts.ContextMaxTokens, HitCount: len(results), ContextCount: len(contextResults), + Policies: buildRetrieverPolicyTraceItems(policies), + } + if ret.TopK <= 0 && len(policies) == 1 { + ret.TopK = policies[0].TopK + } + if ret.ScoreThreshold <= 0 && len(policies) == 1 { + ret.ScoreThreshold = policies[0].ScoreThreshold } if trace != nil { ret.EmbeddingMs = trace.EmbeddingMs @@ -138,3 +209,18 @@ func buildRetrieverTraceSummary(opts KnowledgeRetrieveOptions, contextResults [] } return ret } + +func buildRetrieverPolicyTraceItems(policies []KnowledgeBaseRetrievePolicy) []callbacks.RetrieverPolicyTraceItem { + if len(policies) == 0 { + return nil + } + ret := make([]callbacks.RetrieverPolicyTraceItem, 0, len(policies)) + for _, item := range policies { + ret = append(ret, callbacks.RetrieverPolicyTraceItem{ + KnowledgeBaseID: item.KnowledgeBaseID, + TopK: item.TopK, + ScoreThreshold: item.ScoreThreshold, + }) + } + return ret +}