feat: enhance retriever trace summary by adding policies support and refactoring related methods

This commit is contained in:
mlogclub
2026-04-12 10:59:59 +08:00
parent 7f70614739
commit 0b77d91dfe
3 changed files with 106 additions and 11 deletions
@@ -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) {
@@ -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"`
@@ -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
}