mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-09-01 14:53:07 +08:00
feat(agent): enhance suggested questions API with tag scopes
Update the SuggestedQuestionsRequest to include TagScopes for better filtering of knowledge items based on their parent knowledge bases. Modify related API documentation and frontend interfaces to accommodate the new structure. Ensure backend services handle the new tag scope logic correctly, including tests for the updated functionality.
This commit is contained in:
+16
-6
@@ -307,10 +307,16 @@ type SuggestedQuestion struct {
|
||||
|
||||
// SuggestedQuestionsRequest represents the options for getting suggested questions
|
||||
type SuggestedQuestionsRequest struct {
|
||||
KnowledgeBaseIDs []string // Optional: override agent's KB scope
|
||||
KnowledgeIDs []string // Optional: limit to specific knowledge items
|
||||
TagIDs []string // Optional: limit to knowledge items under these tags
|
||||
Limit int // Optional: max questions to return (default 6)
|
||||
KnowledgeBaseIDs []string // Optional: override agent's KB scope
|
||||
KnowledgeIDs []string // Optional: limit to specific knowledge items
|
||||
TagScopes []SuggestedQuestionTagScope // Optional: limit to tags within their parent KBs
|
||||
Limit int // Optional: max questions to return (default 6)
|
||||
}
|
||||
|
||||
// SuggestedQuestionTagScope preserves the KB-local identity of tag IDs.
|
||||
type SuggestedQuestionTagScope struct {
|
||||
KnowledgeBaseID string `json:"knowledge_base_id"`
|
||||
TagIDs []string `json:"tag_ids"`
|
||||
}
|
||||
|
||||
// SuggestedQuestionsResponse represents the API response for suggested questions
|
||||
@@ -337,8 +343,12 @@ func (c *Client) GetSuggestedQuestions(ctx context.Context, agentID string, requ
|
||||
if len(request.KnowledgeIDs) > 0 {
|
||||
query.Set("knowledge_ids", strings.Join(request.KnowledgeIDs, ","))
|
||||
}
|
||||
if len(request.TagIDs) > 0 {
|
||||
query.Set("tag_ids", strings.Join(request.TagIDs, ","))
|
||||
if len(request.TagScopes) > 0 {
|
||||
encoded, err := json.Marshal(request.TagScopes)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal tag scopes: %w", err)
|
||||
}
|
||||
query.Set("tag_scopes", string(encoded))
|
||||
}
|
||||
if request.Limit > 0 {
|
||||
query.Set("limit", strconv.Itoa(request.Limit))
|
||||
|
||||
@@ -647,6 +647,12 @@ const docTemplate = `{
|
||||
"name": "knowledge_ids",
|
||||
"in": "query"
|
||||
},
|
||||
{
|
||||
"type": "string",
|
||||
"description": "带知识库归属的标签范围(JSON)",
|
||||
"name": "tag_scopes",
|
||||
"in": "query"
|
||||
},
|
||||
{
|
||||
"type": "integer",
|
||||
"description": "返回数量上限(默认6)",
|
||||
|
||||
+7
-1
@@ -640,6 +640,12 @@
|
||||
"name": "knowledge_ids",
|
||||
"in": "query"
|
||||
},
|
||||
{
|
||||
"type": "string",
|
||||
"description": "带知识库归属的标签范围(JSON)",
|
||||
"name": "tag_scopes",
|
||||
"in": "query"
|
||||
},
|
||||
{
|
||||
"type": "integer",
|
||||
"description": "返回数量上限(默认6)",
|
||||
@@ -21784,4 +21790,4 @@
|
||||
"in": "header"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5943,6 +5943,10 @@ paths:
|
||||
in: query
|
||||
name: knowledge_ids
|
||||
type: string
|
||||
- description: 带知识库归属的标签范围(JSON)
|
||||
in: query
|
||||
name: tag_scopes
|
||||
type: string
|
||||
- description: 返回数量上限(默认6)
|
||||
in: query
|
||||
name: limit
|
||||
|
||||
@@ -362,12 +362,17 @@ export interface SuggestedQuestion {
|
||||
// 根据智能体关联的知识库范围返回推荐问题,用于前端对话面板快捷提问
|
||||
export function getSuggestedQuestions(
|
||||
agentId: string,
|
||||
params?: { knowledge_base_ids?: string[]; knowledge_ids?: string[]; tag_ids?: string[]; limit?: number }
|
||||
params?: {
|
||||
knowledge_base_ids?: string[];
|
||||
knowledge_ids?: string[];
|
||||
tag_scopes?: Array<{ knowledge_base_id: string; tag_ids: string[] }>;
|
||||
limit?: number;
|
||||
}
|
||||
) {
|
||||
const query = new URLSearchParams();
|
||||
if (params?.knowledge_base_ids?.length) query.set('knowledge_base_ids', params.knowledge_base_ids.join(','));
|
||||
if (params?.knowledge_ids?.length) query.set('knowledge_ids', params.knowledge_ids.join(','));
|
||||
if (params?.tag_ids?.length) query.set('tag_ids', params.tag_ids.join(','));
|
||||
if (params?.tag_scopes?.length) query.set('tag_scopes', JSON.stringify(params.tag_scopes));
|
||||
if (params?.limit) query.set('limit', String(params.limit));
|
||||
const qs = query.toString();
|
||||
return get<{ data: { questions: SuggestedQuestion[] } }>(`/api/v1/agents/${agentId}/suggested-questions${qs ? '?' + qs : ''}`);
|
||||
|
||||
@@ -470,13 +470,21 @@ export const useSettingsStore = defineStore("settings", {
|
||||
const selectedKBs = this.getSelectedKnowledgeBases();
|
||||
const selectedFiles = this.getSelectedFiles();
|
||||
const tags = this.settings.selectedTags || [];
|
||||
const tagIds = [...new Set(tags.map((t) => t.id).filter(Boolean))];
|
||||
const tagKbIds = [...new Set(tags.map((t) => t.kbId).filter(Boolean))];
|
||||
const kbIds = [...new Set([...selectedKBs, ...tagKbIds])];
|
||||
const tagScopes = Object.entries(tags.reduce<Record<string, string[]>>((scopes, tag) => {
|
||||
if (!tag.id || !tag.kbId) return scopes;
|
||||
(scopes[tag.kbId] ||= []).push(tag.id);
|
||||
return scopes;
|
||||
}, {})).map(([knowledge_base_id, ids]) => ({
|
||||
knowledge_base_id,
|
||||
tag_ids: [...new Set(ids)],
|
||||
}));
|
||||
return {
|
||||
knowledge_base_ids: kbIds.length > 0 ? kbIds : undefined,
|
||||
// A tag's parent KB is only an ownership hint, not an explicit whole-KB
|
||||
// selection. Keep it in tag_scopes so the backend cannot widen a tag to
|
||||
// every document in that KB.
|
||||
knowledge_base_ids: selectedKBs.length > 0 ? selectedKBs : undefined,
|
||||
knowledge_ids: selectedFiles.length > 0 ? selectedFiles : undefined,
|
||||
tag_ids: tagIds.length > 0 ? tagIds : undefined,
|
||||
tag_scopes: tagScopes.length > 0 ? tagScopes : undefined,
|
||||
limit,
|
||||
};
|
||||
},
|
||||
|
||||
@@ -559,14 +559,19 @@ func (t *KnowledgeSearchTool) concurrentSearchByTargets(
|
||||
innerWg.Add(1)
|
||||
go func() {
|
||||
defer innerWg.Done()
|
||||
stVectorThreshold, stKeywordThreshold := st.RecallThresholds(
|
||||
vectorThreshold,
|
||||
keywordThreshold,
|
||||
)
|
||||
searchParams := types.SearchParams{
|
||||
QueryText: q,
|
||||
QueryEmbedding: queryEmbedding,
|
||||
MatchCount: topK,
|
||||
VectorThreshold: vectorThreshold,
|
||||
KeywordThreshold: keywordThreshold,
|
||||
VectorThreshold: stVectorThreshold,
|
||||
KeywordThreshold: stKeywordThreshold,
|
||||
KnowledgeIDs: st.KnowledgeIDs,
|
||||
TagIDs: st.TagIDs,
|
||||
ScopeTagIDs: st.ScopeTagIDs,
|
||||
}
|
||||
kbResults, err := t.knowledgeBaseService.HybridSearch(ctx, st.KnowledgeBaseID, searchParams)
|
||||
if err != nil {
|
||||
@@ -837,7 +842,12 @@ Output only the scores, no explanations or additional text.`,
|
||||
return rankResults[i].RelevanceScore > rankResults[j].RelevanceScore
|
||||
})
|
||||
|
||||
ranked := t.applyModelRerankScores(results, rankResults, t.rerankThreshold())
|
||||
ranked := t.applyModelRerankScores(
|
||||
results,
|
||||
rankResults,
|
||||
t.rerankThreshold(),
|
||||
t.searchTargets.HasRecallThresholdOverride(),
|
||||
)
|
||||
logger.Infof(ctx, "[Tool][KnowledgeSearch] LLM reranked %d/%d results above threshold %.2f",
|
||||
len(ranked), len(results), t.rerankThreshold())
|
||||
return ranked, nil
|
||||
@@ -929,7 +939,12 @@ func (t *KnowledgeSearchTool) rerankWithModel(
|
||||
return nil, fmt.Errorf("rerank call failed: %w", err)
|
||||
}
|
||||
|
||||
ranked := t.applyModelRerankScores(results, rerankResp, t.rerankThreshold())
|
||||
ranked := t.applyModelRerankScores(
|
||||
results,
|
||||
rerankResp,
|
||||
t.rerankThreshold(),
|
||||
t.searchTargets.HasRecallThresholdOverride(),
|
||||
)
|
||||
logger.Infof(
|
||||
ctx,
|
||||
"[Tool][KnowledgeSearch] Reranked %d/%d results above threshold %.2f",
|
||||
@@ -949,7 +964,11 @@ func (t *KnowledgeSearchTool) rerankThreshold() float64 {
|
||||
|
||||
const agentRerankFallbackMinScore = 0.15
|
||||
|
||||
func filterRerankRankResults(rankResults []rerank.RankResult, threshold float64) []rerank.RankResult {
|
||||
func filterRerankRankResults(
|
||||
rankResults []rerank.RankResult,
|
||||
threshold float64,
|
||||
preserveTop bool,
|
||||
) []rerank.RankResult {
|
||||
if len(rankResults) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -966,7 +985,7 @@ func filterRerankRankResults(rankResults []rerank.RankResult, threshold float64)
|
||||
top = r
|
||||
}
|
||||
}
|
||||
if top.RelevanceScore >= agentRerankFallbackMinScore {
|
||||
if preserveTop || top.RelevanceScore >= agentRerankFallbackMinScore {
|
||||
return []rerank.RankResult{top}
|
||||
}
|
||||
}
|
||||
@@ -977,8 +996,9 @@ func (t *KnowledgeSearchTool) applyModelRerankScores(
|
||||
originals []*searchResultWithMeta,
|
||||
rankResults []rerank.RankResult,
|
||||
threshold float64,
|
||||
preserveTop bool,
|
||||
) []*searchResultWithMeta {
|
||||
filtered := filterRerankRankResults(rankResults, threshold)
|
||||
filtered := filterRerankRankResults(rankResults, threshold, preserveTop)
|
||||
out := make([]*searchResultWithMeta, 0, len(filtered))
|
||||
for _, rr := range filtered {
|
||||
if rr.Index < 0 || rr.Index >= len(originals) {
|
||||
|
||||
@@ -14,7 +14,7 @@ func TestFilterRerankRankResults_thresholdAndFallback(t *testing.T) {
|
||||
{Index: 0, RelevanceScore: 0.05},
|
||||
{Index: 1, RelevanceScore: 0.02},
|
||||
}
|
||||
filtered := filterRerankRankResults(rankResults, 0.3)
|
||||
filtered := filterRerankRankResults(rankResults, 0.3, false)
|
||||
if len(filtered) != 0 {
|
||||
t.Fatalf("expected empty filter, got %#v", filtered)
|
||||
}
|
||||
@@ -23,17 +23,26 @@ func TestFilterRerankRankResults_thresholdAndFallback(t *testing.T) {
|
||||
{Index: 0, RelevanceScore: 0.05},
|
||||
{Index: 1, RelevanceScore: 0.20},
|
||||
}
|
||||
filtered = filterRerankRankResults(rankResults, 0.3)
|
||||
filtered = filterRerankRankResults(rankResults, 0.3, false)
|
||||
if len(filtered) != 1 || filtered[0].Index != 1 {
|
||||
t.Fatalf("expected fallback top score, got %#v", filtered)
|
||||
}
|
||||
|
||||
rankResults = []rerank.RankResult{
|
||||
{Index: 0, RelevanceScore: 0.05},
|
||||
{Index: 1, RelevanceScore: 0.02},
|
||||
}
|
||||
filtered = filterRerankRankResults(rankResults, 0.3, true)
|
||||
if len(filtered) != 1 || filtered[0].Index != 0 {
|
||||
t.Fatalf("expected explicit scope to preserve top result, got %#v", filtered)
|
||||
}
|
||||
|
||||
rankResults = []rerank.RankResult{
|
||||
{Index: 0, RelevanceScore: 0.8},
|
||||
{Index: 1, RelevanceScore: 0.4},
|
||||
{Index: 2, RelevanceScore: 0.1},
|
||||
}
|
||||
filtered = filterRerankRankResults(rankResults, 0.3)
|
||||
filtered = filterRerankRankResults(rankResults, 0.3, false)
|
||||
if len(filtered) != 2 {
|
||||
t.Fatalf("expected 2 passing scores, got %#v", filtered)
|
||||
}
|
||||
@@ -59,7 +68,7 @@ func TestApplyModelRerankScores_faqUsesCompositeScale(t *testing.T) {
|
||||
{Index: 0, RelevanceScore: 0.05},
|
||||
{Index: 1, RelevanceScore: 0.9},
|
||||
}
|
||||
out := tool.applyModelRerankScores(originals, rankResults, 0.3)
|
||||
out := tool.applyModelRerankScores(originals, rankResults, 0.3, false)
|
||||
if len(out) != 1 || out[0].ID != "doc-1" {
|
||||
t.Fatalf("weak FAQ should be filtered out, got %#v", out)
|
||||
}
|
||||
|
||||
@@ -957,19 +957,20 @@ func (r *chunkRepository) FAQChunkDiff(
|
||||
}
|
||||
|
||||
// ListRecommendedFAQChunks lists FAQ chunks with the recommended flag set.
|
||||
// Filter by kbIDs and/or knowledgeIDs (OR relationship). At least one must be non-empty.
|
||||
// Filter by explicitly selected kbIDs, knowledgeIDs, and/or FAQ tagIDs (OR relationship).
|
||||
// Returns up to `limit` chunks sorted by updated_at descending.
|
||||
func (r *chunkRepository) ListRecommendedFAQChunks(
|
||||
ctx context.Context,
|
||||
tenantID uint64,
|
||||
kbIDs []string,
|
||||
knowledgeIDs []string,
|
||||
tagIDs []string,
|
||||
limit int,
|
||||
) ([]*types.Chunk, error) {
|
||||
if limit <= 0 {
|
||||
limit = 10
|
||||
}
|
||||
if len(kbIDs) == 0 && len(knowledgeIDs) == 0 {
|
||||
if len(kbIDs) == 0 && len(knowledgeIDs) == 0 && len(tagIDs) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
var chunks []*types.Chunk
|
||||
@@ -977,12 +978,21 @@ func (r *chunkRepository) ListRecommendedFAQChunks(
|
||||
Select("id, knowledge_id, knowledge_base_id, chunk_type, metadata, flags, updated_at").
|
||||
Where("tenant_id = ? AND chunk_type = ? AND status IN ? AND is_enabled = ? AND flags & ? != 0",
|
||||
tenantID, types.ChunkTypeFAQ, []int{int(types.ChunkStatusIndexed), int(types.ChunkStatusDefault)}, true, int(types.ChunkFlagRecommended))
|
||||
if len(knowledgeIDs) > 0 {
|
||||
// 指定了具体知识文档,直接按 knowledge_id 过滤(忽略 kbIDs)
|
||||
query = query.Where("knowledge_id IN ?", knowledgeIDs)
|
||||
} else {
|
||||
query = query.Where("knowledge_base_id IN ?", kbIDs)
|
||||
var scopeClauses []string
|
||||
var scopeArgs []interface{}
|
||||
if len(kbIDs) > 0 {
|
||||
scopeClauses = append(scopeClauses, "knowledge_base_id IN ?")
|
||||
scopeArgs = append(scopeArgs, kbIDs)
|
||||
}
|
||||
if len(knowledgeIDs) > 0 {
|
||||
scopeClauses = append(scopeClauses, "knowledge_id IN ?")
|
||||
scopeArgs = append(scopeArgs, knowledgeIDs)
|
||||
}
|
||||
if len(tagIDs) > 0 {
|
||||
scopeClauses = append(scopeClauses, "tag_id IN ?")
|
||||
scopeArgs = append(scopeArgs, tagIDs)
|
||||
}
|
||||
query = query.Where("("+strings.Join(scopeClauses, " OR ")+")", scopeArgs...)
|
||||
|
||||
orderClause := "RANDOM()"
|
||||
if r.db.Dialector.Name() == "mysql" {
|
||||
|
||||
@@ -215,3 +215,88 @@ func TestUpdateChunk_SQLite_NoNOWError(t *testing.T) {
|
||||
require.NoError(t, db.First(&saved, "id = ?", chunk.ID).Error)
|
||||
assert.Equal(t, "updated content", saved.Content)
|
||||
}
|
||||
|
||||
func makeSuggestedFAQChunk(t *testing.T, kbID, knowledgeID, tagID, question string) *types.Chunk {
|
||||
t.Helper()
|
||||
chunk := makeChunk(kbID, knowledgeID, types.ChunkTypeFAQ)
|
||||
chunk.TagID = tagID
|
||||
chunk.Flags = types.ChunkFlagRecommended
|
||||
require.NoError(t, chunk.SetFAQMetadata(&types.FAQChunkMetadata{StandardQuestion: question}))
|
||||
return chunk
|
||||
}
|
||||
|
||||
func makeSuggestedDocumentChunk(t *testing.T, kbID, knowledgeID, question string) *types.Chunk {
|
||||
t.Helper()
|
||||
chunk := makeChunk(kbID, knowledgeID, types.ChunkTypeText)
|
||||
require.NoError(t, chunk.SetDocumentMetadata(&types.DocumentChunkMetadata{
|
||||
GeneratedQuestions: []types.GeneratedQuestion{{ID: uuid.NewString(), Question: question}},
|
||||
}))
|
||||
return chunk
|
||||
}
|
||||
|
||||
func TestListRecommendedFAQChunks_FiltersByTagWithoutWideningToParentKB(t *testing.T) {
|
||||
db := setupChunkTestDB(t)
|
||||
repo := NewChunkRepository(db)
|
||||
ctx := context.Background()
|
||||
|
||||
selectedTag := uuid.NewString()
|
||||
otherTag := uuid.NewString()
|
||||
selected := makeSuggestedFAQChunk(t, "kb-1", "faq-knowledge", selectedTag, "selected question")
|
||||
other := makeSuggestedFAQChunk(t, "kb-1", "faq-knowledge", otherTag, "other question")
|
||||
require.NoError(t, repo.CreateChunks(ctx, []*types.Chunk{selected, other}))
|
||||
|
||||
got, err := repo.ListRecommendedFAQChunks(ctx, 1, nil, nil, []string{selectedTag}, 10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, got, 1)
|
||||
assert.Equal(t, selected.ID, got[0].ID)
|
||||
}
|
||||
|
||||
func TestListRecommendedFAQChunks_UnionsOnlyExplicitScopes(t *testing.T) {
|
||||
db := setupChunkTestDB(t)
|
||||
repo := NewChunkRepository(db)
|
||||
ctx := context.Background()
|
||||
|
||||
selectedTag := uuid.NewString()
|
||||
tagged := makeSuggestedFAQChunk(t, "kb-tag", "faq-tag", selectedTag, "tagged question")
|
||||
explicitKB := makeSuggestedFAQChunk(t, "kb-explicit", "faq-explicit", uuid.NewString(), "explicit KB question")
|
||||
unselected := makeSuggestedFAQChunk(t, "kb-other", "faq-other", uuid.NewString(), "unselected question")
|
||||
require.NoError(t, repo.CreateChunks(ctx, []*types.Chunk{tagged, explicitKB, unselected}))
|
||||
|
||||
got, err := repo.ListRecommendedFAQChunks(ctx, 1, []string{"kb-explicit"}, nil, []string{selectedTag}, 10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, got, 2)
|
||||
assert.ElementsMatch(t, []string{tagged.ID, explicitKB.ID}, []string{got[0].ID, got[1].ID})
|
||||
}
|
||||
|
||||
func TestListRecentDocumentChunksWithQuestions_KnowledgeScopeDoesNotIncludeSiblingDocuments(t *testing.T) {
|
||||
db := setupChunkTestDB(t)
|
||||
repo := NewChunkRepository(db)
|
||||
ctx := context.Background()
|
||||
|
||||
selected := makeSuggestedDocumentChunk(t, "kb-1", "doc-selected", "selected document question")
|
||||
sibling := makeSuggestedDocumentChunk(t, "kb-1", "doc-sibling", "sibling document question")
|
||||
require.NoError(t, repo.CreateChunks(ctx, []*types.Chunk{selected, sibling}))
|
||||
|
||||
got, err := repo.ListRecentDocumentChunksWithQuestions(ctx, 1, nil, []string{"doc-selected"}, 10)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, got, 1)
|
||||
assert.Equal(t, selected.ID, got[0].ID)
|
||||
}
|
||||
|
||||
func TestListRecentDocumentChunksWithQuestions_UnionsExplicitKBAndKnowledge(t *testing.T) {
|
||||
db := setupChunkTestDB(t)
|
||||
repo := NewChunkRepository(db)
|
||||
ctx := context.Background()
|
||||
|
||||
fromExplicitKB := makeSuggestedDocumentChunk(t, "kb-explicit", "doc-1", "explicit KB question")
|
||||
fromExplicitDocument := makeSuggestedDocumentChunk(t, "kb-other", "doc-selected", "selected document question")
|
||||
unselected := makeSuggestedDocumentChunk(t, "kb-other", "doc-other", "unselected question")
|
||||
require.NoError(t, repo.CreateChunks(ctx, []*types.Chunk{fromExplicitKB, fromExplicitDocument, unselected}))
|
||||
|
||||
got, err := repo.ListRecentDocumentChunksWithQuestions(
|
||||
ctx, 1, []string{"kb-explicit"}, []string{"doc-selected"}, 10,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, got, 2)
|
||||
assert.ElementsMatch(t, []string{fromExplicitKB.ID, fromExplicitDocument.ID}, []string{got[0].ID, got[1].ID})
|
||||
}
|
||||
|
||||
@@ -50,10 +50,11 @@ func agentHasKnowledgeScope(config *types.AgentConfig) bool {
|
||||
if config == nil {
|
||||
return false
|
||||
}
|
||||
if len(config.KnowledgeBases) > 0 || len(config.KnowledgeIDs) > 0 {
|
||||
return true
|
||||
}
|
||||
return len(config.SearchTargets) > 0
|
||||
return types.HasKnowledgeRetrievalScope(
|
||||
config.SearchTargets,
|
||||
config.KnowledgeBases,
|
||||
config.KnowledgeIDs,
|
||||
)
|
||||
}
|
||||
|
||||
// knowledgeBaseIDsForPrompt returns KB IDs to show in runtime_context metadata.
|
||||
|
||||
@@ -233,9 +233,11 @@ func hasKBRetrievalTargets(chatManage *types.ChatManage) bool {
|
||||
if chatManage == nil {
|
||||
return false
|
||||
}
|
||||
return len(chatManage.SearchTargets) > 0 ||
|
||||
len(chatManage.KnowledgeBaseIDs) > 0 ||
|
||||
len(chatManage.KnowledgeIDs) > 0
|
||||
return types.HasKnowledgeRetrievalScope(
|
||||
chatManage.SearchTargets,
|
||||
chatManage.KnowledgeBaseIDs,
|
||||
chatManage.KnowledgeIDs,
|
||||
)
|
||||
}
|
||||
|
||||
func retrievalSearchSource(chatManage *types.ChatManage) string {
|
||||
|
||||
@@ -52,12 +52,17 @@ func (p *PluginSearch) runQueryExpansion(ctx context.Context, chatManage *types.
|
||||
defer wgExp.Done()
|
||||
sem <- struct{}{}
|
||||
defer func() { <-sem }()
|
||||
vectorThreshold, keywordThreshold := t.RecallThresholds(
|
||||
chatManage.VectorThreshold,
|
||||
expKwTh,
|
||||
)
|
||||
paramsExp := types.SearchParams{
|
||||
QueryText: q,
|
||||
VectorThreshold: chatManage.VectorThreshold,
|
||||
KeywordThreshold: expKwTh,
|
||||
VectorThreshold: vectorThreshold,
|
||||
KeywordThreshold: keywordThreshold,
|
||||
MatchCount: expTopK,
|
||||
TagIDs: t.TagIDs,
|
||||
ScopeTagIDs: t.ScopeTagIDs,
|
||||
DisableVectorMatch: false,
|
||||
DisableKeywordsMatch: false,
|
||||
SkipContextEnrichment: true, // Pipeline handles context assembly in merge stage
|
||||
|
||||
@@ -71,19 +71,11 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
return ErrGetRerankModel.WithError(err)
|
||||
}
|
||||
|
||||
// Prepare passages for reranking (excluding DirectLoad results)
|
||||
// Prepare passages for reranking.
|
||||
var passages []string
|
||||
var candidatesToRerank []*types.SearchResult
|
||||
var directLoadResults []*types.SearchResult
|
||||
|
||||
for _, result := range chatManage.SearchResult {
|
||||
if result.MatchType == types.MatchTypeDirectLoad {
|
||||
directLoadResults = append(directLoadResults, result)
|
||||
pipelineInfo(ctx, "Rerank", "direct_load_skip", map[string]interface{}{
|
||||
"chunk_id": result.ID,
|
||||
})
|
||||
continue
|
||||
}
|
||||
passage := getEnrichedPassage(ctx, result)
|
||||
if strings.TrimSpace(passage) == "" {
|
||||
pipelineInfo(ctx, "Rerank", "empty_passage_skip", map[string]interface{}{
|
||||
@@ -99,15 +91,14 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
rerankCtx, rerankSpan := langfuse.GetManager().StartSpan(ctx, langfuse.SpanOptions{
|
||||
Name: "rerank",
|
||||
Input: map[string]interface{}{
|
||||
"query": chatManage.RewriteQuery,
|
||||
"candidate_count": len(candidatesToRerank),
|
||||
"direct_load_count": len(directLoadResults),
|
||||
"rerank_model_id": chatManage.RerankModelID,
|
||||
"threshold": chatManage.RerankThreshold,
|
||||
"rerank_top_k": chatManage.RerankTopK,
|
||||
"faq_priority": chatManage.FAQPriorityEnabled,
|
||||
"faq_score_boost": chatManage.FAQScoreBoost,
|
||||
"passages_preview": passagesPreview,
|
||||
"query": chatManage.RewriteQuery,
|
||||
"candidate_count": len(candidatesToRerank),
|
||||
"rerank_model_id": chatManage.RerankModelID,
|
||||
"threshold": chatManage.RerankThreshold,
|
||||
"rerank_top_k": chatManage.RerankTopK,
|
||||
"faq_priority": chatManage.FAQPriorityEnabled,
|
||||
"faq_score_boost": chatManage.FAQScoreBoost,
|
||||
"passages_preview": passagesPreview,
|
||||
},
|
||||
Metadata: map[string]interface{}{
|
||||
"session_id": chatManage.SessionID,
|
||||
@@ -123,7 +114,6 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
pipelineInfo(ctx, "Rerank", "build_passages", map[string]interface{}{
|
||||
"total_cnt": len(chatManage.SearchResult),
|
||||
"candidate_cnt": len(candidatesToRerank),
|
||||
"direct_cnt": len(directLoadResults),
|
||||
})
|
||||
|
||||
var rerankResp []rerank.RankResult
|
||||
@@ -144,7 +134,7 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
"error": rerankErr.Error(),
|
||||
"candidate_cnt": len(candidatesToRerank),
|
||||
})
|
||||
chatManage.SearchResult = append(directLoadResults, candidatesToRerank...)
|
||||
chatManage.SearchResult = candidatesToRerank
|
||||
spanOutput = map[string]interface{}{
|
||||
"stage": "api_error_fallback",
|
||||
"candidate_count": len(candidatesToRerank),
|
||||
@@ -176,7 +166,7 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
"error": rerankErr.Error(),
|
||||
"candidate_cnt": len(candidatesToRerank),
|
||||
})
|
||||
chatManage.SearchResult = append(directLoadResults, candidatesToRerank...)
|
||||
chatManage.SearchResult = candidatesToRerank
|
||||
spanOutput = map[string]interface{}{
|
||||
"stage": "api_error_fallback",
|
||||
"candidate_count": len(candidatesToRerank),
|
||||
@@ -198,7 +188,7 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
for i := range chatManage.SearchResult {
|
||||
chatManage.SearchResult[i].Metadata = ensureMetadata(chatManage.SearchResult[i].Metadata)
|
||||
}
|
||||
reranked := make([]*types.SearchResult, 0, len(rerankResp)+len(directLoadResults))
|
||||
reranked := make([]*types.SearchResult, 0, len(rerankResp))
|
||||
|
||||
// Process reranked results
|
||||
for _, rr := range rerankResp {
|
||||
@@ -230,16 +220,6 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
reranked = append(reranked, sr)
|
||||
}
|
||||
|
||||
// Process direct load results (bypass rerank model, assume high relevance)
|
||||
for _, sr := range directLoadResults {
|
||||
base := sr.Score
|
||||
sr.Metadata["base_score"] = fmt.Sprintf("%.4f", base)
|
||||
modelScore := 1.0
|
||||
sr.Metadata["model_score"] = fmt.Sprintf("%.4f", modelScore)
|
||||
// Assign high model score for direct load items
|
||||
sr.Score = compositeScore(sr, modelScore, base)
|
||||
reranked = append(reranked, sr)
|
||||
}
|
||||
final := applyMMR(ctx, reranked, chatManage, min(len(reranked), max(1, chatManage.RerankTopK)), 0.7)
|
||||
chatManage.RerankResult = final
|
||||
|
||||
@@ -261,7 +241,6 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
spanOutput = buildRerankSpanOutput(
|
||||
candidatesToRerank,
|
||||
passages,
|
||||
directLoadResults,
|
||||
rawRerankResp,
|
||||
reranked,
|
||||
nil,
|
||||
@@ -274,7 +253,6 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
spanOutput = buildRerankSpanOutput(
|
||||
candidatesToRerank,
|
||||
passages,
|
||||
directLoadResults,
|
||||
rawRerankResp,
|
||||
reranked,
|
||||
chatManage.RerankResult,
|
||||
@@ -290,7 +268,6 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
|
||||
func buildRerankSpanOutput(
|
||||
candidates []*types.SearchResult,
|
||||
passages []string,
|
||||
directLoad []*types.SearchResult,
|
||||
modelScores []rerank.RankResult,
|
||||
composite []*types.SearchResult,
|
||||
final []*types.SearchResult,
|
||||
@@ -319,7 +296,6 @@ func buildRerankSpanOutput(
|
||||
|
||||
out := map[string]interface{}{
|
||||
"candidate_count": len(candidates),
|
||||
"direct_load_count": len(directLoad),
|
||||
"model_result_count": len(modelScores),
|
||||
"composite_count": len(composite),
|
||||
"final_count": len(final),
|
||||
@@ -415,7 +391,7 @@ func (p *PluginRerank) rerank(ctx context.Context,
|
||||
// still has a reasonable score, keep it as a safety net. Skip fallback entirely
|
||||
// when the best score is too low — forcing irrelevant results is worse than
|
||||
// returning nothing and letting the caller handle the empty-result case.
|
||||
const fallbackMinScore = 0.15
|
||||
fallbackMinScore := rerankFallbackMinScore(chatManage.SearchTargets)
|
||||
if len(rankFilter) == 0 && len(rerankResp) > 0 && rerankResp[0].RelevanceScore >= fallbackMinScore {
|
||||
rankFilter = rerankResp[:1]
|
||||
pipelineInfo(ctx, "Rerank", "fallback_top1", map[string]interface{}{
|
||||
@@ -434,6 +410,16 @@ func (p *PluginRerank) rerank(ctx context.Context,
|
||||
return rankFilter, nil
|
||||
}
|
||||
|
||||
func rerankFallbackMinScore(searchTargets types.SearchTargets) float64 {
|
||||
if searchTargets.HasRecallThresholdOverride() {
|
||||
// The user explicitly constrained this turn to a tag/document scope.
|
||||
// Preserve its best candidate instead of letting a global rerank
|
||||
// threshold erase the entire authoritative scope.
|
||||
return 0
|
||||
}
|
||||
return 0.15
|
||||
}
|
||||
|
||||
// ensureMetadata ensures the metadata is not nil
|
||||
func ensureMetadata(m map[string]string) map[string]string {
|
||||
if m == nil {
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
package chatpipeline
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
)
|
||||
|
||||
func TestRerankFallbackMinScoreForExplicitScope(t *testing.T) {
|
||||
if got := rerankFallbackMinScore(nil); got != 0.15 {
|
||||
t.Fatalf("default fallback minimum = %v, want 0.15", got)
|
||||
}
|
||||
|
||||
targets := types.SearchTargets{{DisableRecallThresholds: true}}
|
||||
if got := rerankFallbackMinScore(targets); got != 0 {
|
||||
t.Fatalf("explicit-scope fallback minimum = %v, want 0", got)
|
||||
}
|
||||
}
|
||||
@@ -63,7 +63,11 @@ func (p *PluginSearch) OnEvent(ctx context.Context,
|
||||
eventType types.EventType, chatManage *types.ChatManage, next func() *PluginError,
|
||||
) *PluginError {
|
||||
// Check if we have search targets or web search enabled
|
||||
hasKBTargets := len(chatManage.SearchTargets) > 0 || len(chatManage.KnowledgeBaseIDs) > 0 || len(chatManage.KnowledgeIDs) > 0
|
||||
hasKBTargets := types.HasKnowledgeRetrievalScope(
|
||||
chatManage.SearchTargets,
|
||||
chatManage.KnowledgeBaseIDs,
|
||||
chatManage.KnowledgeIDs,
|
||||
)
|
||||
if !hasKBTargets && !chatManage.WebSearchEnabled {
|
||||
pipelineError(ctx, "Search", "kb_not_found", map[string]interface{}{
|
||||
"session_id": chatManage.SessionID,
|
||||
@@ -456,8 +460,7 @@ func (p *PluginSearch) searchByTargets(
|
||||
return results
|
||||
}
|
||||
|
||||
// searchSingleTarget handles the search logic for a single SearchTarget
|
||||
// with specific knowledge IDs, including direct chunk loading and HybridSearch.
|
||||
// searchSingleTarget performs hybrid retrieval inside one constrained target.
|
||||
func (p *PluginSearch) searchSingleTarget(
|
||||
ctx context.Context,
|
||||
chatManage *types.ChatManage,
|
||||
@@ -467,46 +470,33 @@ func (p *PluginSearch) searchSingleTarget(
|
||||
mu *sync.Mutex,
|
||||
results *[]*types.SearchResult,
|
||||
) {
|
||||
searchKnowledgeIDs := t.KnowledgeIDs
|
||||
|
||||
if t.Type == types.SearchTargetTypeKnowledge && !t.DisableDirectLoad {
|
||||
directResults, skippedIDs := p.tryDirectChunkLoading(ctx, chatManage.TenantID, t.KnowledgeIDs)
|
||||
|
||||
if len(directResults) > 0 {
|
||||
for _, r := range directResults {
|
||||
r.KnowledgeBaseID = t.KnowledgeBaseID
|
||||
}
|
||||
pipelineInfo(ctx, "Search", "direct_load", map[string]interface{}{
|
||||
"kb_id": t.KnowledgeBaseID,
|
||||
"loaded_count": len(directResults),
|
||||
"skipped_ids": len(skippedIDs),
|
||||
})
|
||||
mu.Lock()
|
||||
*results = append(*results, directResults...)
|
||||
mu.Unlock()
|
||||
}
|
||||
|
||||
if len(skippedIDs) == 0 && len(t.KnowledgeIDs) > 0 {
|
||||
return
|
||||
}
|
||||
searchKnowledgeIDs = skippedIDs
|
||||
}
|
||||
|
||||
if t.Type == types.SearchTargetTypeKnowledge && len(searchKnowledgeIDs) == 0 {
|
||||
if t.Type == types.SearchTargetTypeKnowledge && len(t.KnowledgeIDs) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
vectorThreshold, keywordThreshold := t.RecallThresholds(
|
||||
chatManage.VectorThreshold,
|
||||
chatManage.KeywordThreshold,
|
||||
)
|
||||
if t.DisableRecallThresholds {
|
||||
pipelineInfo(ctx, "Search", "explicit_scope_threshold_override", map[string]interface{}{
|
||||
"kb_id": t.KnowledgeBaseID,
|
||||
"knowledge_id_count": len(t.KnowledgeIDs),
|
||||
"tag_id_count": len(t.TagIDs),
|
||||
})
|
||||
}
|
||||
params := types.SearchParams{
|
||||
QueryText: queryText,
|
||||
QueryEmbedding: queryEmbedding,
|
||||
VectorThreshold: chatManage.VectorThreshold,
|
||||
KeywordThreshold: chatManage.KeywordThreshold,
|
||||
VectorThreshold: vectorThreshold,
|
||||
KeywordThreshold: keywordThreshold,
|
||||
MatchCount: chatManage.EmbeddingTopK,
|
||||
TagIDs: t.TagIDs,
|
||||
ScopeTagIDs: t.ScopeTagIDs,
|
||||
SkipContextEnrichment: true,
|
||||
}
|
||||
if t.Type == types.SearchTargetTypeKnowledge {
|
||||
params.KnowledgeIDs = searchKnowledgeIDs
|
||||
params.KnowledgeIDs = t.KnowledgeIDs
|
||||
}
|
||||
res, err := p.knowledgeBaseService.HybridSearch(ctx, t.KnowledgeBaseID, params)
|
||||
if err != nil {
|
||||
@@ -528,96 +518,6 @@ func (p *PluginSearch) searchSingleTarget(
|
||||
mu.Unlock()
|
||||
}
|
||||
|
||||
// tryDirectChunkLoading attempts to load chunks for given knowledge IDs directly
|
||||
// Returns loaded results and a list of knowledge IDs that were skipped (e.g. due to size limits)
|
||||
func (p *PluginSearch) tryDirectChunkLoading(ctx context.Context, tenantID uint64, knowledgeIDs []string) ([]*types.SearchResult, []string) {
|
||||
if len(knowledgeIDs) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// Limit direct loading to avoid OOM or context overflow
|
||||
// 50 chunks * ~500 chars/chunk ~= 25k chars
|
||||
const maxTotalChunks = 50
|
||||
|
||||
var allChunks []*types.Chunk
|
||||
var skippedIDs []string
|
||||
loadedKnowledgeIDs := make(map[string]bool)
|
||||
|
||||
for _, kid := range knowledgeIDs {
|
||||
// Optimization: Check chunk count first if possible?
|
||||
chunks, err := p.chunkService.ListChunksByKnowledgeID(ctx, kid)
|
||||
if err != nil {
|
||||
logger.Warnf(ctx, "DirectLoad: Failed to list chunks for knowledge %s: %v", kid, err)
|
||||
skippedIDs = append(skippedIDs, kid)
|
||||
continue
|
||||
}
|
||||
|
||||
if len(allChunks)+len(chunks) > maxTotalChunks {
|
||||
logger.Infof(ctx, "DirectLoad: Skipped knowledge %s due to size limit (%d + %d > %d)",
|
||||
kid, len(allChunks), len(chunks), maxTotalChunks)
|
||||
skippedIDs = append(skippedIDs, kid)
|
||||
continue
|
||||
}
|
||||
allChunks = append(allChunks, chunks...)
|
||||
loadedKnowledgeIDs[kid] = true
|
||||
}
|
||||
|
||||
if len(allChunks) == 0 {
|
||||
return nil, skippedIDs
|
||||
}
|
||||
|
||||
// Fetch Knowledge metadata
|
||||
var uniqueKIDs []string
|
||||
for kid := range loadedKnowledgeIDs {
|
||||
uniqueKIDs = append(uniqueKIDs, kid)
|
||||
}
|
||||
|
||||
knowledgeMap := make(map[string]*types.Knowledge)
|
||||
if len(uniqueKIDs) > 0 {
|
||||
knowledges, err := p.knowledgeService.GetKnowledgeBatchWithSharedAccess(ctx, tenantID, uniqueKIDs)
|
||||
if err != nil {
|
||||
logger.Warnf(ctx, "DirectLoad: Failed to fetch knowledge batch: %v", err)
|
||||
// Continue without metadata
|
||||
} else {
|
||||
for _, k := range knowledges {
|
||||
knowledgeMap[k.ID] = k
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var results []*types.SearchResult
|
||||
for _, chunk := range allChunks {
|
||||
res := &types.SearchResult{
|
||||
ID: chunk.ID,
|
||||
Content: chunk.Content,
|
||||
Score: 1.0, // Maximum score for direct matches
|
||||
KnowledgeID: chunk.KnowledgeID,
|
||||
ChunkIndex: chunk.ChunkIndex,
|
||||
MatchType: types.MatchTypeDirectLoad,
|
||||
ChunkType: string(chunk.ChunkType),
|
||||
ParentChunkID: chunk.ParentChunkID,
|
||||
ImageInfo: chunk.ImageInfo,
|
||||
ChunkMetadata: chunk.Metadata,
|
||||
StartAt: chunk.StartAt,
|
||||
EndAt: chunk.EndAt,
|
||||
}
|
||||
|
||||
if k, ok := knowledgeMap[chunk.KnowledgeID]; ok {
|
||||
res.KnowledgeTitle = k.Title
|
||||
res.KnowledgeFilename = k.FileName
|
||||
res.KnowledgeSource = k.Source
|
||||
res.KnowledgeChannel = k.Channel
|
||||
res.Metadata = k.GetMetadata()
|
||||
}
|
||||
|
||||
results = append(results, res)
|
||||
}
|
||||
|
||||
searchutil.EnrichSearchResultsImageInfo(ctx, p.chunkService.GetRepository(), tenantID, results)
|
||||
|
||||
return results, skippedIDs
|
||||
}
|
||||
|
||||
// searchWebIfEnabled executes web search when enabled and returns converted results
|
||||
func (p *PluginSearch) searchWebIfEnabled(ctx context.Context, chatManage *types.ChatManage) []*types.SearchResult {
|
||||
if !chatManage.WebSearchEnabled || p.webSearchService == nil || p.tenantService == nil {
|
||||
|
||||
@@ -471,10 +471,10 @@ func (s *customAgentService) GetSuggestedQuestions(
|
||||
agentID string,
|
||||
kbIDs []string,
|
||||
knowledgeIDs []string,
|
||||
tagIDs []string,
|
||||
tagScopes []types.TagScope,
|
||||
limit int,
|
||||
) ([]types.SuggestedQuestion, error) {
|
||||
return s.getSuggestedQuestions(ctx, agentID, kbIDs, knowledgeIDs, tagIDs, limit, true)
|
||||
return s.getSuggestedQuestions(ctx, agentID, kbIDs, knowledgeIDs, tagScopes, limit, true)
|
||||
}
|
||||
|
||||
func (s *customAgentService) GetKnowledgeSuggestedQuestions(
|
||||
@@ -482,10 +482,10 @@ func (s *customAgentService) GetKnowledgeSuggestedQuestions(
|
||||
agentID string,
|
||||
kbIDs []string,
|
||||
knowledgeIDs []string,
|
||||
tagIDs []string,
|
||||
tagScopes []types.TagScope,
|
||||
limit int,
|
||||
) ([]types.SuggestedQuestion, error) {
|
||||
return s.getSuggestedQuestions(ctx, agentID, kbIDs, knowledgeIDs, tagIDs, limit, false)
|
||||
return s.getSuggestedQuestions(ctx, agentID, kbIDs, knowledgeIDs, tagScopes, limit, false)
|
||||
}
|
||||
|
||||
func (s *customAgentService) getSuggestedQuestions(
|
||||
@@ -493,7 +493,7 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
agentID string,
|
||||
kbIDs []string,
|
||||
knowledgeIDs []string,
|
||||
tagIDs []string,
|
||||
tagScopes []types.TagScope,
|
||||
limit int,
|
||||
includeCurated bool,
|
||||
) ([]types.SuggestedQuestion, error) {
|
||||
@@ -504,7 +504,8 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
if err := types.AuthorizeTenantAPIKeyKnowledgeTargets(ctx, kbIDs, knowledgeIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := types.AuthorizeTenantAPIKeyOptionalTagIDs(ctx, tagIDs); err != nil {
|
||||
scopeTagIDs := flattenTagScopeIDs(tagScopes)
|
||||
if err := types.AuthorizeTenantAPIKeyOptionalTagIDs(ctx, scopeTagIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -520,7 +521,8 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var result []types.SuggestedQuestion
|
||||
var curated []types.SuggestedQuestion
|
||||
starterMode := types.SuggestionModeKnowledge
|
||||
|
||||
if includeCurated {
|
||||
suggestionConfig := agent.Config.QuestionSuggestions
|
||||
@@ -530,6 +532,7 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
if limit > suggestionConfig.Starters.Count {
|
||||
limit = suggestionConfig.Starters.Count
|
||||
}
|
||||
starterMode = suggestionConfig.Starters.Mode
|
||||
// Add curated agent prompts first (highest priority).
|
||||
if suggestionConfig.Starters.Mode == types.SuggestionModeCurated ||
|
||||
suggestionConfig.Starters.Mode == types.SuggestionModeHybrid {
|
||||
@@ -537,35 +540,37 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
if strings.TrimSpace(prompt) == "" {
|
||||
continue
|
||||
}
|
||||
result = append(result, types.SuggestedQuestion{
|
||||
curated = append(curated, types.SuggestedQuestion{
|
||||
Question: prompt,
|
||||
Source: "agent_config",
|
||||
})
|
||||
}
|
||||
}
|
||||
if suggestionConfig.Starters.Mode == types.SuggestionModeCurated {
|
||||
return s.truncateQuestions(result, limit), nil
|
||||
return s.truncateQuestions(curated, limit), nil
|
||||
}
|
||||
}
|
||||
|
||||
if len(tagIDs) > 0 {
|
||||
resolved, err := s.resolveKnowledgeIDsFromTags(ctx, tenantID, tagIDs)
|
||||
resolvedTags := resolvedSuggestionTagScopes{}
|
||||
if len(scopeTagIDs) > 0 {
|
||||
var err error
|
||||
resolvedTags, err = s.resolveSuggestionTagScopes(ctx, tenantID, tagScopes)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{
|
||||
"agent_id": agentID,
|
||||
"tag_ids": tagIDs,
|
||||
"agent_id": agentID,
|
||||
"scope_tag_ids": scopeTagIDs,
|
||||
})
|
||||
return s.truncateQuestions(result, limit), nil
|
||||
return finalizeStarterSuggestions(curated, nil, starterMode, limit), nil
|
||||
}
|
||||
knowledgeIDs = mergeUniqueStrings(knowledgeIDs, resolved)
|
||||
if len(knowledgeIDs) == 0 {
|
||||
return s.truncateQuestions(result, limit), nil
|
||||
knowledgeIDs = mergeUniqueStrings(knowledgeIDs, resolvedTags.KnowledgeIDs)
|
||||
if len(knowledgeIDs) == 0 && len(resolvedTags.TagIDsByTenant) == 0 {
|
||||
return finalizeStarterSuggestions(curated, nil, starterMode, limit), nil
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Determine knowledge base scope
|
||||
effectiveKBIDs := kbIDs
|
||||
if len(effectiveKBIDs) == 0 && len(knowledgeIDs) == 0 {
|
||||
if len(effectiveKBIDs) == 0 && len(knowledgeIDs) == 0 && len(resolvedTags.TagIDsByTenant) == 0 {
|
||||
// Use agent's KB configuration
|
||||
switch agent.Config.KBSelectionMode {
|
||||
case "all":
|
||||
@@ -575,7 +580,7 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
"agent_id": agentID,
|
||||
})
|
||||
// Return what we have so far (agent_config suggestions)
|
||||
return s.truncateQuestions(result, limit), nil
|
||||
return finalizeStarterSuggestions(curated, nil, starterMode, limit), nil
|
||||
}
|
||||
// Honor the agent's implicit/explicit capability requirements so
|
||||
// e.g. a quick-answer (RAG-only) agent doesn't surface wiki-only
|
||||
@@ -593,12 +598,16 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
effectiveKBIDs = agent.Config.KnowledgeBases
|
||||
case "none":
|
||||
// No KB access, return agent_config suggestions only
|
||||
return s.truncateQuestions(result, limit), nil
|
||||
return finalizeStarterSuggestions(curated, nil, starterMode, limit), nil
|
||||
default:
|
||||
// Default to agent's configured KBs
|
||||
effectiveKBIDs = agent.Config.KnowledgeBases
|
||||
}
|
||||
}
|
||||
// Match the chat retrieval target semantics: a tag scope narrows its parent
|
||||
// KB even when that KB is present in the agent's preselected KB list. Other
|
||||
// explicitly selected KBs remain additive.
|
||||
effectiveKBIDs = excludeSuggestionStrings(effectiveKBIDs, resolvedTags.KnowledgeBaseIDs)
|
||||
|
||||
filteredKBIDs, err := types.FilterKnowledgeBasesForTenantAPIKeyScope(ctx, kbIDs, effectiveKBIDs)
|
||||
if err != nil {
|
||||
@@ -606,20 +615,17 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
}
|
||||
effectiveKBIDs = filteredKBIDs
|
||||
|
||||
if len(effectiveKBIDs) == 0 && len(knowledgeIDs) == 0 {
|
||||
return s.truncateQuestions(result, limit), nil
|
||||
if len(effectiveKBIDs) == 0 && len(knowledgeIDs) == 0 && len(resolvedTags.TagIDsByTenant) == 0 {
|
||||
return finalizeStarterSuggestions(curated, nil, starterMode, limit), nil
|
||||
}
|
||||
|
||||
// Deduplicate questions we've already collected
|
||||
seen := make(map[string]bool)
|
||||
for _, q := range result {
|
||||
for _, q := range curated {
|
||||
seen[q.Question] = true
|
||||
}
|
||||
|
||||
remaining := limit - len(result)
|
||||
if remaining <= 0 {
|
||||
return s.truncateQuestions(result, limit), nil
|
||||
}
|
||||
remaining := limit
|
||||
|
||||
// 3. Collect candidate chunks from both FAQ and Document KBs,
|
||||
// grouped by knowledge_id for diversity.
|
||||
@@ -641,16 +647,21 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
// rows live under that tenant. Without this grouping a caller in tenant A
|
||||
// querying a KB shared from tenant B would hit `tenant_id = A` and get zero
|
||||
// rows back — the symptom is "suggested questions never appear for shared KBs".
|
||||
kbGroups := s.groupKBIDsByEffectiveTenant(ctx, tenantID, queryKBIDs)
|
||||
scopeKBIDs := mergeUniqueStrings(queryKBIDs, resolvedTags.KnowledgeBaseIDs)
|
||||
kbGroups := s.groupKBIDsByEffectiveTenant(ctx, tenantID, scopeKBIDs)
|
||||
// Always keep the caller's tenant in the iteration so knowledge_ids-only
|
||||
// requests (no kbIDs) still execute one query under the caller's tenant.
|
||||
if len(queryKBIDs) == 0 {
|
||||
if len(scopeKBIDs) == 0 {
|
||||
kbGroups[tenantID] = nil
|
||||
}
|
||||
|
||||
// Collect FAQ recommended chunks
|
||||
for groupTenantID, groupKBIDs := range kbGroups {
|
||||
faqChunks, err := s.chunkRepo.ListRecommendedFAQChunks(ctx, groupTenantID, groupKBIDs, queryKnowledgeIDs, fetchLimit)
|
||||
explicitGroupKBIDs := intersectSuggestionStrings(groupKBIDs, queryKBIDs)
|
||||
groupTagIDs := resolvedTags.TagIDsByTenant[groupTenantID]
|
||||
faqChunks, err := s.chunkRepo.ListRecommendedFAQChunks(
|
||||
ctx, groupTenantID, explicitGroupKBIDs, queryKnowledgeIDs, groupTagIDs, fetchLimit,
|
||||
)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{
|
||||
"agent_id": agentID,
|
||||
@@ -677,7 +688,8 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
|
||||
// Collect Document chunks with generated questions
|
||||
for groupTenantID, groupKBIDs := range kbGroups {
|
||||
docChunks, err := s.chunkRepo.ListRecentDocumentChunksWithQuestions(ctx, groupTenantID, groupKBIDs, queryKnowledgeIDs, fetchLimit)
|
||||
explicitGroupKBIDs := intersectSuggestionStrings(groupKBIDs, queryKBIDs)
|
||||
docChunks, err := s.chunkRepo.ListRecentDocumentChunksWithQuestions(ctx, groupTenantID, explicitGroupKBIDs, queryKnowledgeIDs, fetchLimit)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{
|
||||
"agent_id": agentID,
|
||||
@@ -703,11 +715,10 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
}
|
||||
}
|
||||
|
||||
// Collect Wiki pages as a fallback source. This covers Wiki-only KBs where no
|
||||
// document chunks carry AI-generated questions (question_generation is skipped
|
||||
// when the KB does not need an embedding model). knowledge_id filter is
|
||||
// intentionally ignored here because wiki pages are authored at the KB level
|
||||
// and are not 1:1 with source knowledge items.
|
||||
// Collect Wiki pages as a fallback source, but only for KBs the caller selected
|
||||
// explicitly. A tag's parent KB is merely an ownership boundary; widening a
|
||||
// tag-only scope to arbitrary Wiki pages would make the suggestions unanswerable
|
||||
// inside the user's selected range.
|
||||
//
|
||||
// Skip entirely for quick-answer (RAG-only) agents: those can't ever
|
||||
// retrieve a wiki page, so surfacing wiki-derived suggestions would lure
|
||||
@@ -715,10 +726,11 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
// context. Smart-reasoning agents that opt in to wiki tools keep this.
|
||||
if agent.Config.AgentMode != types.AgentModeQuickAnswer && s.wikiPageRepo != nil {
|
||||
for groupTenantID, groupKBIDs := range kbGroups {
|
||||
if len(groupKBIDs) == 0 {
|
||||
explicitGroupKBIDs := intersectSuggestionStrings(groupKBIDs, queryKBIDs)
|
||||
if len(explicitGroupKBIDs) == 0 {
|
||||
continue
|
||||
}
|
||||
wikiPages, err := s.wikiPageRepo.ListRecentForSuggestions(ctx, groupTenantID, groupKBIDs, fetchLimit)
|
||||
wikiPages, err := s.wikiPageRepo.ListRecentForSuggestions(ctx, groupTenantID, explicitGroupKBIDs, fetchLimit)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{
|
||||
"agent_id": agentID,
|
||||
@@ -757,17 +769,18 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
})
|
||||
|
||||
// Round-robin pick one question from each document in turn.
|
||||
knowledgeResult := make([]types.SuggestedQuestion, 0, limit)
|
||||
offsets := make(map[string]int, len(bucketKeys))
|
||||
for len(result) < limit {
|
||||
for len(knowledgeResult) < limit {
|
||||
picked := false
|
||||
for _, key := range bucketKeys {
|
||||
if len(result) >= limit {
|
||||
if len(knowledgeResult) >= limit {
|
||||
break
|
||||
}
|
||||
qs := buckets[key]
|
||||
idx := offsets[key]
|
||||
if idx < len(qs) {
|
||||
result = append(result, qs[idx])
|
||||
knowledgeResult = append(knowledgeResult, qs[idx])
|
||||
offsets[key] = idx + 1
|
||||
picked = true
|
||||
}
|
||||
@@ -777,52 +790,188 @@ func (s *customAgentService) getSuggestedQuestions(
|
||||
}
|
||||
}
|
||||
|
||||
return s.truncateQuestions(result, limit), nil
|
||||
return finalizeStarterSuggestions(curated, knowledgeResult, starterMode, limit), nil
|
||||
}
|
||||
|
||||
func (s *customAgentService) resolveKnowledgeIDsFromTags(
|
||||
type resolvedSuggestionTagScopes struct {
|
||||
KnowledgeBaseIDs []string
|
||||
KnowledgeIDs []string
|
||||
TagIDsByTenant map[uint64][]string
|
||||
}
|
||||
|
||||
// resolveSuggestionTagScopes keeps tag ownership separate from whole-KB
|
||||
// selection. Document tags become concrete knowledge IDs; FAQ tags remain
|
||||
// chunk tag filters. Scoped inputs also let shared-KB tags resolve against the
|
||||
// source tenant that owns the tag and chunk rows.
|
||||
func (s *customAgentService) resolveSuggestionTagScopes(
|
||||
ctx context.Context,
|
||||
tenantID uint64,
|
||||
tagIDs []string,
|
||||
) ([]string, error) {
|
||||
if len(tagIDs) == 0 || s.tagRepo == nil || s.knowledgeRepo == nil {
|
||||
return nil, nil
|
||||
}
|
||||
tags, err := s.tagRepo.GetByIDs(ctx, tenantID, tagIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(tags) == 0 {
|
||||
return nil, nil
|
||||
callerTenantID uint64,
|
||||
tagScopes []types.TagScope,
|
||||
) (resolvedSuggestionTagScopes, error) {
|
||||
result := resolvedSuggestionTagScopes{TagIDsByTenant: make(map[uint64][]string)}
|
||||
if len(tagScopes) == 0 || s.tagRepo == nil || s.knowledgeRepo == nil || s.kbService == nil {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
byKB := make(map[string][]string)
|
||||
for _, tag := range tags {
|
||||
byKB[tag.KnowledgeBaseID] = append(byKB[tag.KnowledgeBaseID], tag.ID)
|
||||
for _, scope := range tagScopes {
|
||||
if scope.KnowledgeBaseID == "" {
|
||||
continue
|
||||
}
|
||||
for _, tagID := range scope.TagIDs {
|
||||
if tagID == "" {
|
||||
continue
|
||||
}
|
||||
byKB[scope.KnowledgeBaseID] = append(byKB[scope.KnowledgeBaseID], tagID)
|
||||
}
|
||||
}
|
||||
return mergeKnowledgeIDsFromTagGroups(ctx, s.knowledgeRepo, tenantID, byKB)
|
||||
if len(byKB) == 0 {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
kbIDs := make([]string, 0, len(byKB))
|
||||
for kbID := range byKB {
|
||||
kbIDs = append(kbIDs, kbID)
|
||||
}
|
||||
kbGroups := s.groupKBIDsByEffectiveTenant(ctx, callerTenantID, kbIDs)
|
||||
for tenantID, groupKBIDs := range kbGroups {
|
||||
for _, kbID := range groupKBIDs {
|
||||
requested := mergeUniqueStrings(nil, byKB[kbID])
|
||||
tags, err := s.tagRepo.GetByIDs(ctx, tenantID, requested)
|
||||
if err != nil {
|
||||
return result, err
|
||||
}
|
||||
requestedSet := make(map[string]bool, len(requested))
|
||||
for _, id := range requested {
|
||||
requestedSet[id] = true
|
||||
}
|
||||
validTagIDs := make([]string, 0, len(tags))
|
||||
for _, tag := range tags {
|
||||
if tag != nil && tag.KnowledgeBaseID == kbID && requestedSet[tag.ID] {
|
||||
validTagIDs = append(validTagIDs, tag.ID)
|
||||
}
|
||||
}
|
||||
if len(validTagIDs) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
result.KnowledgeBaseIDs = mergeUniqueStrings(result.KnowledgeBaseIDs, []string{kbID})
|
||||
result.TagIDsByTenant[tenantID] = mergeUniqueStrings(result.TagIDsByTenant[tenantID], validTagIDs)
|
||||
knowledgeIDs, err := s.knowledgeRepo.ListIDsByTagIDs(ctx, tenantID, kbID, validTagIDs)
|
||||
if err != nil {
|
||||
return result, err
|
||||
}
|
||||
result.KnowledgeIDs = mergeUniqueStrings(result.KnowledgeIDs, knowledgeIDs)
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func mergeKnowledgeIDsFromTagGroups(
|
||||
ctx context.Context,
|
||||
knowledgeRepo interfaces.KnowledgeRepository,
|
||||
tenantID uint64,
|
||||
byKB map[string][]string,
|
||||
) ([]string, error) {
|
||||
seen := make(map[string]bool)
|
||||
var out []string
|
||||
for kbID, ids := range byKB {
|
||||
kids, err := knowledgeRepo.ListIDsByTagIDs(ctx, tenantID, kbID, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, kid := range kids {
|
||||
if !seen[kid] {
|
||||
seen[kid] = true
|
||||
out = append(out, kid)
|
||||
}
|
||||
func flattenTagScopeIDs(scopes []types.TagScope) []string {
|
||||
var ids []string
|
||||
for _, scope := range scopes {
|
||||
ids = mergeUniqueStrings(ids, scope.TagIDs)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func intersectSuggestionStrings(values, allowed []string) []string {
|
||||
if len(values) == 0 || len(allowed) == 0 {
|
||||
return nil
|
||||
}
|
||||
allowedSet := make(map[string]bool, len(allowed))
|
||||
for _, value := range allowed {
|
||||
allowedSet[value] = true
|
||||
}
|
||||
var result []string
|
||||
for _, value := range values {
|
||||
if value != "" && allowedSet[value] {
|
||||
result = append(result, value)
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
return result
|
||||
}
|
||||
|
||||
func excludeSuggestionStrings(values, excluded []string) []string {
|
||||
if len(values) == 0 || len(excluded) == 0 {
|
||||
return values
|
||||
}
|
||||
excludedSet := make(map[string]bool, len(excluded))
|
||||
for _, value := range excluded {
|
||||
excludedSet[value] = true
|
||||
}
|
||||
result := make([]string, 0, len(values))
|
||||
for _, value := range values {
|
||||
if value != "" && !excludedSet[value] {
|
||||
result = append(result, value)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func finalizeStarterSuggestions(
|
||||
curated []types.SuggestedQuestion,
|
||||
knowledge []types.SuggestedQuestion,
|
||||
mode string,
|
||||
limit int,
|
||||
) []types.SuggestedQuestion {
|
||||
if limit <= 0 {
|
||||
return []types.SuggestedQuestion{}
|
||||
}
|
||||
switch mode {
|
||||
case types.SuggestionModeCurated:
|
||||
return truncateSuggestedQuestions(curated, limit)
|
||||
case types.SuggestionModeHybrid:
|
||||
return mergeHybridStarterSuggestions(curated, knowledge, limit)
|
||||
default:
|
||||
return truncateSuggestedQuestions(knowledge, limit)
|
||||
}
|
||||
}
|
||||
|
||||
// mergeHybridStarterSuggestions prioritizes curated starters while reserving
|
||||
// about one third of visible slots for scope-aware knowledge questions.
|
||||
func mergeHybridStarterSuggestions(
|
||||
curated []types.SuggestedQuestion,
|
||||
knowledge []types.SuggestedQuestion,
|
||||
limit int,
|
||||
) []types.SuggestedQuestion {
|
||||
if limit <= 0 {
|
||||
return []types.SuggestedQuestion{}
|
||||
}
|
||||
knowledgeSlots := 0
|
||||
if limit > 1 {
|
||||
knowledgeSlots = (limit + 1) / 3
|
||||
}
|
||||
curatedSlots := limit - knowledgeSlots
|
||||
result := make([]types.SuggestedQuestion, 0, limit)
|
||||
seen := make(map[string]bool, limit)
|
||||
appendFrom := func(items []types.SuggestedQuestion, max int) {
|
||||
added := 0
|
||||
for _, item := range items {
|
||||
if len(result) == limit || (max >= 0 && added == max) {
|
||||
return
|
||||
}
|
||||
key := strings.ToLower(strings.TrimSpace(item.Question))
|
||||
if key == "" || seen[key] {
|
||||
continue
|
||||
}
|
||||
seen[key] = true
|
||||
result = append(result, item)
|
||||
added++
|
||||
}
|
||||
}
|
||||
appendFrom(curated, curatedSlots)
|
||||
appendFrom(knowledge, knowledgeSlots)
|
||||
appendFrom(curated, -1)
|
||||
appendFrom(knowledge, -1)
|
||||
return result
|
||||
}
|
||||
|
||||
func truncateSuggestedQuestions(questions []types.SuggestedQuestion, limit int) []types.SuggestedQuestion {
|
||||
if len(questions) > limit {
|
||||
return questions[:limit]
|
||||
}
|
||||
return questions
|
||||
}
|
||||
|
||||
func mergeUniqueStrings(base, extra []string) []string {
|
||||
|
||||
@@ -33,15 +33,22 @@ func TestGetSuggestedQuestionsRejectsKnowledgeIDsForRestrictedKey(t *testing.T)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetSuggestedQuestionsRejectsTagIDsForRestrictedKey(t *testing.T) {
|
||||
func TestGetSuggestedQuestionsRejectsTagScopesForRestrictedKey(t *testing.T) {
|
||||
ctx := types.WithTenantAPIKeyScope(context.Background(), types.TenantAPIKeyScope{
|
||||
KnowledgeBaseIDs: types.StringArray{"kb-1"},
|
||||
})
|
||||
ctx = context.WithValue(ctx, types.TenantIDContextKey, uint64(1))
|
||||
|
||||
svc := &customAgentService{}
|
||||
_, err := svc.GetSuggestedQuestions(ctx, "agent-1", nil, nil, []string{"tag-1"}, 6)
|
||||
_, err := svc.GetSuggestedQuestions(
|
||||
ctx,
|
||||
"agent-1",
|
||||
nil,
|
||||
nil,
|
||||
[]types.TagScope{{KnowledgeBaseID: "kb-1", TagIDs: []string{"tag-1"}}},
|
||||
6,
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatal("expected forbidden for tag_ids under KB-restricted key")
|
||||
t.Fatal("expected forbidden for tag_scopes under KB-restricted key")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type suggestionTagRepo struct {
|
||||
interfaces.KnowledgeTagRepository
|
||||
tagsByTenant map[uint64][]*types.KnowledgeTag
|
||||
}
|
||||
|
||||
func (r *suggestionTagRepo) GetByIDs(_ context.Context, tenantID uint64, ids []string) ([]*types.KnowledgeTag, error) {
|
||||
wanted := make(map[string]bool, len(ids))
|
||||
for _, id := range ids {
|
||||
wanted[id] = true
|
||||
}
|
||||
var result []*types.KnowledgeTag
|
||||
for _, tag := range r.tagsByTenant[tenantID] {
|
||||
if tag != nil && wanted[tag.ID] {
|
||||
result = append(result, tag)
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
type suggestionKnowledgeRepo struct {
|
||||
interfaces.KnowledgeRepository
|
||||
idsByTenantAndKB map[uint64]map[string][]string
|
||||
}
|
||||
|
||||
func (r *suggestionKnowledgeRepo) ListIDsByTagIDs(
|
||||
_ context.Context,
|
||||
tenantID uint64,
|
||||
kbID string,
|
||||
_ []string,
|
||||
) ([]string, error) {
|
||||
return append([]string(nil), r.idsByTenantAndKB[tenantID][kbID]...), nil
|
||||
}
|
||||
|
||||
type suggestionKBService struct {
|
||||
interfaces.KnowledgeBaseService
|
||||
kbs map[string]*types.KnowledgeBase
|
||||
}
|
||||
|
||||
func (s *suggestionKBService) GetKnowledgeBasesByIDsOnly(
|
||||
_ context.Context,
|
||||
ids []string,
|
||||
) ([]*types.KnowledgeBase, error) {
|
||||
result := make([]*types.KnowledgeBase, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if kb := s.kbs[id]; kb != nil {
|
||||
result = append(result, kb)
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
type suggestionKBShareService struct {
|
||||
interfaces.KBShareService
|
||||
allowed map[string]bool
|
||||
}
|
||||
|
||||
func (s *suggestionKBShareService) HasTenantKBPermission(
|
||||
_ context.Context,
|
||||
kbID string,
|
||||
_ uint64,
|
||||
_ types.TenantRole,
|
||||
_ types.OrgMemberRole,
|
||||
) (bool, error) {
|
||||
return s.allowed[kbID], nil
|
||||
}
|
||||
|
||||
func TestResolveSuggestionTagScopes_UsesSourceTenantForSharedKB(t *testing.T) {
|
||||
const (
|
||||
callerTenant = uint64(1)
|
||||
sourceTenant = uint64(2)
|
||||
kbID = "shared-kb"
|
||||
tagID = "shared-tag"
|
||||
)
|
||||
svc := &customAgentService{
|
||||
tagRepo: &suggestionTagRepo{tagsByTenant: map[uint64][]*types.KnowledgeTag{
|
||||
sourceTenant: {{ID: tagID, TenantID: sourceTenant, KnowledgeBaseID: kbID}},
|
||||
}},
|
||||
knowledgeRepo: &suggestionKnowledgeRepo{idsByTenantAndKB: map[uint64]map[string][]string{
|
||||
sourceTenant: {kbID: {"doc-in-tag"}},
|
||||
}},
|
||||
kbService: &suggestionKBService{kbs: map[string]*types.KnowledgeBase{
|
||||
kbID: {ID: kbID, TenantID: sourceTenant},
|
||||
}},
|
||||
kbShareService: &suggestionKBShareService{allowed: map[string]bool{kbID: true}},
|
||||
}
|
||||
|
||||
resolved, err := svc.resolveSuggestionTagScopes(
|
||||
context.Background(),
|
||||
callerTenant,
|
||||
[]types.TagScope{{KnowledgeBaseID: kbID, TagIDs: []string{tagID}}},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{kbID}, resolved.KnowledgeBaseIDs)
|
||||
assert.Equal(t, []string{"doc-in-tag"}, resolved.KnowledgeIDs)
|
||||
assert.Equal(t, []string{tagID}, resolved.TagIDsByTenant[sourceTenant])
|
||||
assert.Empty(t, resolved.TagIDsByTenant[callerTenant])
|
||||
}
|
||||
|
||||
func TestMergeHybridStarterSuggestions_ReservesKnowledgeSlots(t *testing.T) {
|
||||
curated := []types.SuggestedQuestion{
|
||||
{Question: "curated 1", Source: "agent_config"},
|
||||
{Question: "curated 2", Source: "agent_config"},
|
||||
{Question: "curated 3", Source: "agent_config"},
|
||||
{Question: "curated 4", Source: "agent_config"},
|
||||
{Question: "curated 5", Source: "agent_config"},
|
||||
{Question: "curated 6", Source: "agent_config"},
|
||||
}
|
||||
knowledge := []types.SuggestedQuestion{
|
||||
{Question: "knowledge 1", Source: "document"},
|
||||
{Question: "knowledge 2", Source: "faq"},
|
||||
{Question: "knowledge 3", Source: "document"},
|
||||
}
|
||||
|
||||
got := mergeHybridStarterSuggestions(curated, knowledge, 6)
|
||||
require.Len(t, got, 6)
|
||||
assert.Equal(t, []string{
|
||||
"curated 1", "curated 2", "curated 3", "curated 4", "knowledge 1", "knowledge 2",
|
||||
}, []string{got[0].Question, got[1].Question, got[2].Question, got[3].Question, got[4].Question, got[5].Question})
|
||||
}
|
||||
|
||||
func TestMergeHybridStarterSuggestions_BackfillsWhenKnowledgeIsEmpty(t *testing.T) {
|
||||
curated := []types.SuggestedQuestion{
|
||||
{Question: "curated 1"}, {Question: "curated 2"}, {Question: "curated 3"},
|
||||
}
|
||||
got := mergeHybridStarterSuggestions(curated, nil, 3)
|
||||
require.Len(t, got, 3)
|
||||
assert.Equal(t, []string{"curated 1", "curated 2", "curated 3"}, []string{
|
||||
got[0].Question, got[1].Question, got[2].Question,
|
||||
})
|
||||
}
|
||||
|
||||
func TestExcludeSuggestionStrings_TagScopeOverridesSameKnowledgeBase(t *testing.T) {
|
||||
got := excludeSuggestionStrings([]string{"kb-with-tag", "kb-explicit"}, []string{"kb-with-tag"})
|
||||
assert.Equal(t, []string{"kb-explicit"}, got)
|
||||
}
|
||||
@@ -195,6 +195,7 @@ func (s *knowledgeBaseService) HybridSearch(ctx context.Context,
|
||||
"kb_ids": searchKBIDs,
|
||||
"knowledge_ids": params.KnowledgeIDs,
|
||||
"tag_ids": params.TagIDs,
|
||||
"scope_tag_ids": params.ScopeTagIDs,
|
||||
"match_count": matchCount,
|
||||
"vector_threshold": params.VectorThreshold,
|
||||
"keyword_threshold": params.KeywordThreshold,
|
||||
@@ -203,9 +204,9 @@ func (s *knowledgeBaseService) HybridSearch(ctx context.Context,
|
||||
"group_count": len(groups),
|
||||
},
|
||||
Metadata: map[string]interface{}{
|
||||
"primary_kb_id": kb.ID,
|
||||
"primary_kb_type": string(kb.Type),
|
||||
"embedding_model_id": kb.EmbeddingModelID,
|
||||
"primary_kb_id": kb.ID,
|
||||
"primary_kb_type": string(kb.Type),
|
||||
"embedding_model_id": kb.EmbeddingModelID,
|
||||
"has_query_embedding": len(params.QueryEmbedding) > 0,
|
||||
},
|
||||
})
|
||||
|
||||
@@ -425,7 +425,7 @@ func (s *messageSuggestionService) generateFromKnowledge(
|
||||
message.AgentID,
|
||||
message.ExecutionContext.KnowledgeBaseIDs,
|
||||
knowledgeIDs,
|
||||
message.ExecutionContext.TagIDs,
|
||||
message.ExecutionContext.TagScopes,
|
||||
poolSize,
|
||||
)
|
||||
if err != nil {
|
||||
@@ -439,7 +439,7 @@ func (s *messageSuggestionService) generateFromKnowledge(
|
||||
message.AgentID,
|
||||
message.ExecutionContext.KnowledgeBaseIDs,
|
||||
message.ExecutionContext.KnowledgeIDs,
|
||||
message.ExecutionContext.TagIDs,
|
||||
message.ExecutionContext.TagScopes,
|
||||
poolSize,
|
||||
)
|
||||
if err != nil {
|
||||
|
||||
@@ -308,6 +308,18 @@ func (s *sessionService) buildAgentConfig(
|
||||
return nil, fmt.Errorf("build search targets: %w", err)
|
||||
}
|
||||
agentConfig.SearchTargets = searchTargets
|
||||
// Document tags are stored in knowledge_tag_relations, so document-KB tag
|
||||
// scopes are resolved to concrete knowledge IDs before retrieval. Preserve
|
||||
// those resolved IDs as this turn's pinned documents as well: otherwise the
|
||||
// Agent tools are correctly constrained behind the scenes, but the model only
|
||||
// sees a bound KB and does not know which documents the user explicitly chose.
|
||||
if len(req.TagScopes) > 0 {
|
||||
agentConfig.KnowledgeIDs = mergeResolvedTagKnowledgeIDs(
|
||||
agentConfig.KnowledgeIDs,
|
||||
searchTargets,
|
||||
req.TagScopes,
|
||||
)
|
||||
}
|
||||
logger.Infof(ctx, "Agent search targets built: %d targets", len(searchTargets))
|
||||
|
||||
if agentConfig.MaxContextTokens <= 0 {
|
||||
@@ -317,6 +329,31 @@ func (s *sessionService) buildAgentConfig(
|
||||
return agentConfig, nil
|
||||
}
|
||||
|
||||
func mergeResolvedTagKnowledgeIDs(
|
||||
existing []string,
|
||||
searchTargets types.SearchTargets,
|
||||
tagScopes []types.TagScope,
|
||||
) []string {
|
||||
tagKBs := make(map[string]bool, len(tagScopes))
|
||||
for _, scope := range tagScopes {
|
||||
if scope.KnowledgeBaseID != "" && len(scope.TagIDs) > 0 {
|
||||
tagKBs[scope.KnowledgeBaseID] = true
|
||||
}
|
||||
}
|
||||
if len(tagKBs) == 0 {
|
||||
return uniqueNonEmptyStrings(existing)
|
||||
}
|
||||
|
||||
merged := append([]string(nil), existing...)
|
||||
for _, target := range searchTargets {
|
||||
if target == nil || !tagKBs[target.KnowledgeBaseID] || target.Type != types.SearchTargetTypeKnowledge {
|
||||
continue
|
||||
}
|
||||
merged = append(merged, target.KnowledgeIDs...)
|
||||
}
|
||||
return uniqueNonEmptyStrings(merged)
|
||||
}
|
||||
|
||||
// applyPerRequestSkillScope narrows the agent's skill whitelist to the @Skill
|
||||
// mentions for this turn and records the pinned set for the <must_use> hint.
|
||||
// It is a no-op when no skills were mentioned or skills are disabled.
|
||||
|
||||
@@ -158,8 +158,11 @@ func (s *sessionService) KnowledgeQA(
|
||||
// rewrite, fallback, FAQ strategy, history turns)
|
||||
s.applyAgentOverridesToChatManage(ctx, req.CustomAgent, chatManage)
|
||||
|
||||
// Determine pipeline based on knowledge bases availability and web search setting
|
||||
hasKB := len(knowledgeBaseIDs) > 0 || len(knowledgeIDs) > 0
|
||||
// Determine pipeline based on the effective knowledge retrieval scope and
|
||||
// web search setting. Tag-only mentions leave the raw KB/knowledge ID slices
|
||||
// empty but produce SearchTargets, so the unified targets must participate in
|
||||
// this decision or the request is incorrectly downgraded to pure chat.
|
||||
hasKB := types.HasKnowledgeRetrievalScope(searchTargets, knowledgeBaseIDs, knowledgeIDs)
|
||||
needsRAG := hasKB || req.WebSearchEnabled
|
||||
hasHistory := chatManage.MaxRounds > 0
|
||||
|
||||
@@ -541,10 +544,11 @@ func (s *sessionService) buildSearchTargets(
|
||||
kbTenant = tenantID // fallback
|
||||
}
|
||||
targets = append(targets, &types.SearchTarget{
|
||||
Type: types.SearchTargetTypeKnowledge,
|
||||
KnowledgeBaseID: kbID,
|
||||
TenantID: kbTenant,
|
||||
KnowledgeIDs: kidList,
|
||||
Type: types.SearchTargetTypeKnowledge,
|
||||
KnowledgeBaseID: kbID,
|
||||
TenantID: kbTenant,
|
||||
KnowledgeIDs: kidList,
|
||||
DisableRecallThresholds: true,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -574,25 +578,28 @@ func (s *sessionService) buildSearchTargets(
|
||||
continue
|
||||
}
|
||||
targets = append(targets, &types.SearchTarget{
|
||||
Type: types.SearchTargetTypeKnowledge,
|
||||
KnowledgeBaseID: kbID,
|
||||
TenantID: kbTenant,
|
||||
KnowledgeIDs: tagKnowledgeIDs,
|
||||
DisableDirectLoad: true,
|
||||
Type: types.SearchTargetTypeKnowledge,
|
||||
KnowledgeBaseID: kbID,
|
||||
TenantID: kbTenant,
|
||||
KnowledgeIDs: tagKnowledgeIDs,
|
||||
ScopeTagIDs: append([]string(nil), tagIDs...),
|
||||
DisableRecallThresholds: true,
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
target := &types.SearchTarget{
|
||||
Type: types.SearchTargetTypeKnowledgeBase,
|
||||
KnowledgeBaseID: kbID,
|
||||
TenantID: kbTenant,
|
||||
TagIDs: append([]string(nil), tagIDs...),
|
||||
Type: types.SearchTargetTypeKnowledgeBase,
|
||||
KnowledgeBaseID: kbID,
|
||||
TenantID: kbTenant,
|
||||
TagIDs: append([]string(nil), tagIDs...),
|
||||
ScopeTagIDs: append([]string(nil), tagIDs...),
|
||||
DisableRecallThresholds: true,
|
||||
}
|
||||
if len(explicitKnowledgeIDs) > 0 {
|
||||
target.Type = types.SearchTargetTypeKnowledge
|
||||
target.KnowledgeIDs = explicitKnowledgeIDs
|
||||
target.DisableDirectLoad = true
|
||||
target.DisableRecallThresholds = true
|
||||
}
|
||||
targets = append(targets, target)
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/config"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -89,6 +90,7 @@ func knowledgeBelongsToKB(knowledges []*types.Knowledge, knowledgeID string, kbI
|
||||
|
||||
func newTagTargetSessionService() *sessionService {
|
||||
return &sessionService{
|
||||
cfg: &config.Config{},
|
||||
knowledgeBaseService: &tagTargetKnowledgeBaseService{
|
||||
kbs: map[string]*types.KnowledgeBase{
|
||||
"doc-kb": {ID: "doc-kb", TenantID: 100, Type: types.KnowledgeBaseTypeDocument},
|
||||
@@ -110,6 +112,41 @@ func newTagTargetSessionService() *sessionService {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildAgentConfig_TagOnlyScopePreservesRetrievalTarget(t *testing.T) {
|
||||
svc := newTagTargetSessionService()
|
||||
agent := &types.CustomAgent{
|
||||
ID: "agent-1",
|
||||
TenantID: 100,
|
||||
Config: types.CustomAgentConfig{
|
||||
AgentMode: types.AgentModeSmartReasoning,
|
||||
KBSelectionMode: "all",
|
||||
WebSearchProviderID: "provider-1",
|
||||
},
|
||||
}
|
||||
req := &types.QARequest{
|
||||
Session: &types.Session{ID: "session-1", TenantID: 100},
|
||||
CustomAgent: agent,
|
||||
TagScopes: []types.TagScope{
|
||||
{KnowledgeBaseID: "doc-kb", TagIDs: []string{"tag-a"}},
|
||||
},
|
||||
}
|
||||
|
||||
agentConfig, err := svc.buildAgentConfig(
|
||||
tagTargetContext(),
|
||||
req,
|
||||
&types.Tenant{ID: 100},
|
||||
100,
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, agentConfig.KnowledgeBases)
|
||||
require.Len(t, agentConfig.SearchTargets, 1)
|
||||
assert.Equal(t, types.SearchTargetTypeKnowledge, agentConfig.SearchTargets[0].Type)
|
||||
assert.ElementsMatch(t, []string{"doc-1", "doc-3"}, agentConfig.SearchTargets[0].KnowledgeIDs)
|
||||
assert.ElementsMatch(t, []string{"doc-1", "doc-3"}, agentConfig.KnowledgeIDs)
|
||||
assert.True(t, agentHasKnowledgeScope(agentConfig))
|
||||
}
|
||||
|
||||
func tagTargetContext() context.Context {
|
||||
return context.WithValue(context.Background(), types.TenantIDContextKey, uint64(100))
|
||||
}
|
||||
@@ -131,7 +168,26 @@ func TestBuildSearchTargets_DocumentTagScopeResolvesKnowledgeIDs(t *testing.T) {
|
||||
assert.Equal(t, "doc-kb", targets[0].KnowledgeBaseID)
|
||||
assert.ElementsMatch(t, []string{"doc-1", "doc-3"}, targets[0].KnowledgeIDs)
|
||||
assert.Empty(t, targets[0].TagIDs)
|
||||
assert.True(t, targets[0].DisableDirectLoad)
|
||||
assert.ElementsMatch(t, []string{"tag-a"}, targets[0].ScopeTagIDs)
|
||||
assert.True(t, targets[0].DisableRecallThresholds)
|
||||
}
|
||||
|
||||
func TestBuildSearchTargets_ExplicitKnowledgeScopeDisablesRecallThresholds(t *testing.T) {
|
||||
svc := newTagTargetSessionService()
|
||||
|
||||
targets, err := svc.buildSearchTargets(
|
||||
tagTargetContext(),
|
||||
100,
|
||||
nil,
|
||||
[]string{"doc-1"},
|
||||
nil,
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Len(t, targets, 1)
|
||||
assert.Equal(t, types.SearchTargetTypeKnowledge, targets[0].Type)
|
||||
assert.Equal(t, []string{"doc-1"}, targets[0].KnowledgeIDs)
|
||||
assert.True(t, targets[0].DisableRecallThresholds)
|
||||
}
|
||||
|
||||
func TestBuildSearchTargets_DocumentTagScopeIntersectsExplicitKnowledgeIDs(t *testing.T) {
|
||||
@@ -149,7 +205,8 @@ func TestBuildSearchTargets_DocumentTagScopeIntersectsExplicitKnowledgeIDs(t *te
|
||||
require.Len(t, targets, 1)
|
||||
assert.Equal(t, types.SearchTargetTypeKnowledge, targets[0].Type)
|
||||
assert.Equal(t, []string{"doc-3"}, targets[0].KnowledgeIDs)
|
||||
assert.True(t, targets[0].DisableDirectLoad)
|
||||
assert.ElementsMatch(t, []string{"tag-a"}, targets[0].ScopeTagIDs)
|
||||
assert.True(t, targets[0].DisableRecallThresholds)
|
||||
}
|
||||
|
||||
func TestBuildSearchTargets_FAQTagScopeKeepsIndexTagFilter(t *testing.T) {
|
||||
@@ -168,7 +225,8 @@ func TestBuildSearchTargets_FAQTagScopeKeepsIndexTagFilter(t *testing.T) {
|
||||
assert.Equal(t, types.SearchTargetTypeKnowledgeBase, targets[0].Type)
|
||||
assert.Equal(t, "faq-kb", targets[0].KnowledgeBaseID)
|
||||
assert.ElementsMatch(t, []string{"tag-a", "tag-b"}, targets[0].TagIDs)
|
||||
assert.False(t, targets[0].DisableDirectLoad)
|
||||
assert.ElementsMatch(t, []string{"tag-a", "tag-b"}, targets[0].ScopeTagIDs)
|
||||
assert.True(t, targets[0].DisableRecallThresholds)
|
||||
}
|
||||
|
||||
func TestBuildSearchTargets_FullKBWithTagScopeSkipsFullKBTarget(t *testing.T) {
|
||||
@@ -215,7 +273,24 @@ func TestBuildSearchTargets_DocumentTagScopeWithMissingKBMetadata(t *testing.T)
|
||||
require.Len(t, targets, 1)
|
||||
assert.Equal(t, types.SearchTargetTypeKnowledge, targets[0].Type)
|
||||
assert.ElementsMatch(t, []string{"doc-1", "doc-3"}, targets[0].KnowledgeIDs)
|
||||
assert.True(t, targets[0].DisableDirectLoad)
|
||||
assert.True(t, targets[0].DisableRecallThresholds)
|
||||
}
|
||||
|
||||
func TestMergeResolvedTagKnowledgeIDs_OnlyIncludesTagScopedTargets(t *testing.T) {
|
||||
got := mergeResolvedTagKnowledgeIDs(
|
||||
[]string{"existing-doc"},
|
||||
types.SearchTargets{
|
||||
{Type: types.SearchTargetTypeKnowledge, KnowledgeBaseID: "tag-kb", KnowledgeIDs: []string{"tag-doc-1", "tag-doc-2"}},
|
||||
{Type: types.SearchTargetTypeKnowledge, KnowledgeBaseID: "other-kb", KnowledgeIDs: []string{"other-doc"}},
|
||||
{Type: types.SearchTargetTypeKnowledgeBase, KnowledgeBaseID: "faq-kb", TagIDs: []string{"faq-tag"}},
|
||||
},
|
||||
[]types.TagScope{
|
||||
{KnowledgeBaseID: "tag-kb", TagIDs: []string{"tag-a"}},
|
||||
{KnowledgeBaseID: "faq-kb", TagIDs: []string{"faq-tag"}},
|
||||
},
|
||||
)
|
||||
|
||||
assert.ElementsMatch(t, []string{"existing-doc", "tag-doc-1", "tag-doc-2"}, got)
|
||||
}
|
||||
|
||||
type tagTargetKnowledgeServiceWithError struct {
|
||||
|
||||
@@ -2,6 +2,7 @@ package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -561,6 +562,7 @@ func (h *CustomAgentHandler) GetAgentTypePresets(c *gin.Context) {
|
||||
// @Param id path string true "智能体ID"
|
||||
// @Param knowledge_base_ids query string false "知识库ID列表(逗号分隔),覆盖智能体默认配置"
|
||||
// @Param knowledge_ids query string false "知识ID列表(逗号分隔),限定到具体文档"
|
||||
// @Param tag_scopes query string false "带知识库归属的标签范围(JSON)"
|
||||
// @Param limit query int false "返回数量上限(默认6)"
|
||||
// @Success 200 {object} map[string]interface{} "推荐问题列表"
|
||||
// @Failure 400 {object} errors.AppError "请求参数错误"
|
||||
@@ -598,12 +600,11 @@ func (h *CustomAgentHandler) GetSuggestedQuestions(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
var tagIDs []string
|
||||
if tagIDsStr := strings.TrimSpace(c.Query("tag_ids")); tagIDsStr != "" {
|
||||
for _, id := range strings.Split(tagIDsStr, ",") {
|
||||
if trimmed := strings.TrimSpace(id); trimmed != "" {
|
||||
tagIDs = append(tagIDs, trimmed)
|
||||
}
|
||||
var tagScopes []types.TagScope
|
||||
if raw := strings.TrimSpace(c.Query("tag_scopes")); raw != "" {
|
||||
if err := json.Unmarshal([]byte(raw), &tagScopes); err != nil {
|
||||
c.Error(errors.NewBadRequestError("tag_scopes must be valid JSON"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
@@ -614,10 +615,10 @@ func (h *CustomAgentHandler) GetSuggestedQuestions(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Getting suggested questions for agent %s, kbIDs: %v, tagIDs: %v, limit: %d",
|
||||
secutils.SanitizeForLog(id), kbIDs, tagIDs, limit)
|
||||
logger.Infof(ctx, "Getting suggested questions for agent %s, kbIDs: %v, tagScopes: %d, limit: %d",
|
||||
secutils.SanitizeForLog(id), kbIDs, len(tagScopes), limit)
|
||||
|
||||
questions, err := h.service.GetSuggestedQuestions(ctx, id, kbIDs, knowledgeIDs, tagIDs, limit)
|
||||
questions, err := h.service.GetSuggestedQuestions(ctx, id, kbIDs, knowledgeIDs, tagScopes, limit)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{
|
||||
"agent_id": id,
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"testing"
|
||||
|
||||
apperrors "github.com/Tencent/WeKnora/internal/errors"
|
||||
@@ -11,21 +12,24 @@ import (
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type suggestedQuestionsAgentService struct {
|
||||
interfaces.CustomAgentService
|
||||
err error
|
||||
err error
|
||||
tagScopes []types.TagScope
|
||||
}
|
||||
|
||||
func (s *suggestedQuestionsAgentService) GetSuggestedQuestions(
|
||||
context.Context,
|
||||
string,
|
||||
[]string,
|
||||
[]string,
|
||||
[]string,
|
||||
int,
|
||||
_ context.Context,
|
||||
_ string,
|
||||
_ []string,
|
||||
_ []string,
|
||||
tagScopes []types.TagScope,
|
||||
_ int,
|
||||
) ([]types.SuggestedQuestion, error) {
|
||||
s.tagScopes = tagScopes
|
||||
return nil, s.err
|
||||
}
|
||||
|
||||
@@ -47,3 +51,44 @@ func TestGetSuggestedQuestionsPreservesAppErrorStatus(t *testing.T) {
|
||||
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusForbidden, rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetSuggestedQuestionsParsesScopedTags(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.Use(middleware.ErrorHandler())
|
||||
|
||||
service := &suggestedQuestionsAgentService{}
|
||||
h := &CustomAgentHandler{service: service}
|
||||
r.GET("/agents/:id/suggested-questions", h.GetSuggestedQuestions)
|
||||
|
||||
rawScopes := `[{"knowledge_base_id":"kb-1","tag_ids":["tag-1","tag-2"]}]`
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(
|
||||
http.MethodGet,
|
||||
"/agents/agent-1/suggested-questions?tag_scopes="+url.QueryEscape(rawScopes),
|
||||
nil,
|
||||
)
|
||||
r.ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
|
||||
require.Equal(t, []types.TagScope{{KnowledgeBaseID: "kb-1", TagIDs: []string{"tag-1", "tag-2"}}}, service.tagScopes)
|
||||
}
|
||||
|
||||
func TestGetSuggestedQuestionsRejectsInvalidScopedTags(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.Use(middleware.ErrorHandler())
|
||||
|
||||
h := &CustomAgentHandler{service: &suggestedQuestionsAgentService{}}
|
||||
r.GET("/agents/:id/suggested-questions", h.GetSuggestedQuestions)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(
|
||||
http.MethodGet,
|
||||
"/agents/agent-1/suggested-questions?tag_scopes="+url.QueryEscape("not-json"),
|
||||
nil,
|
||||
)
|
||||
r.ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusBadRequest, rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
@@ -314,6 +314,7 @@ func (h *Handler) parseQARequest(c *gin.Context, logPrefix string) (*qaRequestCo
|
||||
secutils.SanitizeForLogArray(kbIDs),
|
||||
secutils.SanitizeForLogArray(knowledgeIDs),
|
||||
secutils.SanitizeForLogArray(tagIDs),
|
||||
tagScopes,
|
||||
secutils.SanitizeForLogArray(mcpServiceIDs),
|
||||
secutils.SanitizeForLogArray(skillNames),
|
||||
request.WebSearchEnabled,
|
||||
@@ -372,6 +373,7 @@ func buildMessageExecutionContext(
|
||||
knowledgeBaseIDs []string,
|
||||
knowledgeIDs []string,
|
||||
tagIDs []string,
|
||||
tagScopes []types.TagScope,
|
||||
mcpServiceIDs []string,
|
||||
skillNames []string,
|
||||
webSearchEnabled bool,
|
||||
@@ -385,6 +387,7 @@ func buildMessageExecutionContext(
|
||||
KnowledgeBaseIDs: knowledgeBaseIDs,
|
||||
KnowledgeIDs: knowledgeIDs,
|
||||
TagIDs: tagIDs,
|
||||
TagScopes: cloneTagScopes(tagScopes),
|
||||
MCPServiceIDs: mcpServiceIDs,
|
||||
SkillNames: skillNames,
|
||||
WebSearchEnabled: webSearchEnabled,
|
||||
@@ -418,12 +421,14 @@ func buildMessageExecutionContext(
|
||||
KnowledgeBaseIDs []string `json:"knowledge_base_ids,omitempty"`
|
||||
KnowledgeIDs []string `json:"knowledge_ids,omitempty"`
|
||||
TagIDs []string `json:"tag_ids,omitempty"`
|
||||
TagScopes []types.TagScope `json:"tag_scopes,omitempty"`
|
||||
ModelID string `json:"model_id,omitempty"`
|
||||
}{
|
||||
QuestionSuggestions: snapshot.QuestionSuggestions,
|
||||
KnowledgeBaseIDs: knowledgeBaseIDs,
|
||||
KnowledgeIDs: knowledgeIDs,
|
||||
TagIDs: tagIDs,
|
||||
TagScopes: snapshot.TagScopes,
|
||||
ModelID: modelID,
|
||||
}
|
||||
if encoded, err := json.Marshal(hashInput); err == nil {
|
||||
@@ -434,6 +439,23 @@ func buildMessageExecutionContext(
|
||||
return snapshot, agent.ID, agentTenantID, modelID
|
||||
}
|
||||
|
||||
func cloneTagScopes(scopes []types.TagScope) []types.TagScope {
|
||||
if len(scopes) == 0 {
|
||||
return nil
|
||||
}
|
||||
cloned := make([]types.TagScope, 0, len(scopes))
|
||||
for _, scope := range scopes {
|
||||
if scope.KnowledgeBaseID == "" || len(scope.TagIDs) == 0 {
|
||||
continue
|
||||
}
|
||||
cloned = append(cloned, types.TagScope{
|
||||
KnowledgeBaseID: scope.KnowledgeBaseID,
|
||||
TagIDs: append([]string(nil), scope.TagIDs...),
|
||||
})
|
||||
}
|
||||
return cloned
|
||||
}
|
||||
|
||||
// resolveEnableMemory decides whether the memory pipeline runs for this
|
||||
// request. See the call-site comment in parseQARequest for the resolution
|
||||
// order. Lookup errors are logged but never propagate — a failure to read
|
||||
|
||||
@@ -175,13 +175,16 @@ func (c *ChatManage) Clone() *ChatManage {
|
||||
copy(kidsCopy, t.KnowledgeIDs)
|
||||
tagIDsCopy := make([]string, len(t.TagIDs))
|
||||
copy(tagIDsCopy, t.TagIDs)
|
||||
scopeTagIDsCopy := make([]string, len(t.ScopeTagIDs))
|
||||
copy(scopeTagIDsCopy, t.ScopeTagIDs)
|
||||
searchTargets[i] = &SearchTarget{
|
||||
Type: t.Type,
|
||||
KnowledgeBaseID: t.KnowledgeBaseID,
|
||||
TenantID: t.TenantID,
|
||||
KnowledgeIDs: kidsCopy,
|
||||
TagIDs: tagIDsCopy,
|
||||
DisableDirectLoad: t.DisableDirectLoad,
|
||||
Type: t.Type,
|
||||
KnowledgeBaseID: t.KnowledgeBaseID,
|
||||
TenantID: t.TenantID,
|
||||
KnowledgeIDs: kidsCopy,
|
||||
TagIDs: tagIDsCopy,
|
||||
ScopeTagIDs: scopeTagIDsCopy,
|
||||
DisableRecallThresholds: t.DisableRecallThresholds,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,7 +21,7 @@ const (
|
||||
MatchTypeRelationChunk // 关系Chunk匹配类型
|
||||
MatchTypeGraph
|
||||
MatchTypeWebSearch // 网络搜索匹配类型
|
||||
MatchTypeDirectLoad // 直接加载匹配类型
|
||||
MatchTypeDirectLoad // Deprecated: reserved to preserve serialized enum values
|
||||
MatchTypeDataAnalysis // 数据分析匹配类型
|
||||
)
|
||||
|
||||
|
||||
@@ -99,9 +99,9 @@ type ChunkRepository interface {
|
||||
FAQChunkDiff(ctx context.Context, srcTenantID uint64, srcKBID string, dstTenantID uint64, dstKBID string) (chunksToAdd []string, chunksToDelete []string, err error)
|
||||
|
||||
// ListRecommendedFAQChunks lists FAQ chunks with the recommended flag set.
|
||||
// Filter by kbIDs and/or knowledgeIDs. At least one of them must be non-empty.
|
||||
// Filter by explicitly selected kbIDs, knowledgeIDs, and/or FAQ tagIDs.
|
||||
// Returns up to `limit` chunks sorted by updated_at descending.
|
||||
ListRecommendedFAQChunks(ctx context.Context, tenantID uint64, kbIDs []string, knowledgeIDs []string, limit int) ([]*types.Chunk, error)
|
||||
ListRecommendedFAQChunks(ctx context.Context, tenantID uint64, kbIDs []string, knowledgeIDs []string, tagIDs []string, limit int) ([]*types.Chunk, error)
|
||||
|
||||
// ListRecentDocumentChunksWithQuestions lists recent document chunks that have generated questions.
|
||||
// Filter by kbIDs and/or knowledgeIDs. At least one of them must be non-empty.
|
||||
|
||||
@@ -73,17 +73,17 @@ type CustomAgentService interface {
|
||||
// - agentID: Agent ID
|
||||
// - kbIDs: Optional knowledge base IDs to override agent config
|
||||
// - knowledgeIDs: Optional knowledge item IDs to further filter
|
||||
// - tagIDs: Optional knowledge tag IDs; resolved to knowledge item IDs (OR semantics)
|
||||
// - tagScopes: Optional KB-scoped knowledge tags (OR semantics within each KB)
|
||||
// - limit: Maximum number of questions to return
|
||||
// Returns:
|
||||
// - List of suggested questions
|
||||
// - Possible errors
|
||||
GetSuggestedQuestions(ctx context.Context, agentID string, kbIDs []string, knowledgeIDs []string, tagIDs []string, limit int) ([]types.SuggestedQuestion, error)
|
||||
GetSuggestedQuestions(ctx context.Context, agentID string, kbIDs []string, knowledgeIDs []string, tagScopes []types.TagScope, limit int) ([]types.SuggestedQuestion, error)
|
||||
|
||||
// GetKnowledgeSuggestedQuestions returns only knowledge-derived candidates.
|
||||
// It is independent of whether starter suggestions are enabled and is used
|
||||
// as a source/fallback for contextual follow-up generation.
|
||||
GetKnowledgeSuggestedQuestions(ctx context.Context, agentID string, kbIDs []string, knowledgeIDs []string, tagIDs []string, limit int) ([]types.SuggestedQuestion, error)
|
||||
GetKnowledgeSuggestedQuestions(ctx context.Context, agentID string, kbIDs []string, knowledgeIDs []string, tagScopes []types.TagScope, limit int) ([]types.SuggestedQuestion, error)
|
||||
}
|
||||
|
||||
// CustomAgentRepository defines the custom agent repository interface
|
||||
|
||||
@@ -268,6 +268,7 @@ type MessageExecutionContext struct {
|
||||
KnowledgeBaseIDs []string `json:"knowledge_base_ids,omitempty"`
|
||||
KnowledgeIDs []string `json:"knowledge_ids,omitempty"`
|
||||
TagIDs []string `json:"tag_ids,omitempty"`
|
||||
TagScopes []TagScope `json:"tag_scopes,omitempty"`
|
||||
MCPServiceIDs []string `json:"mcp_service_ids,omitempty"`
|
||||
SkillNames []string `json:"skill_names,omitempty"`
|
||||
WebSearchEnabled bool `json:"web_search_enabled"`
|
||||
|
||||
@@ -36,16 +36,50 @@ type SearchTarget struct {
|
||||
KnowledgeIDs []string `json:"knowledge_ids,omitempty"`
|
||||
// TagIDs limits retrieval to chunks/documents carrying any of these KB-local tags.
|
||||
TagIDs []string `json:"tag_ids,omitempty"`
|
||||
// DisableDirectLoad forces the target through retrieval even when it is
|
||||
// represented as specific knowledge IDs. Tag-derived document scopes need
|
||||
// this so tag filtering limits the candidate documents without loading every
|
||||
// matching document chunk as context.
|
||||
DisableDirectLoad bool `json:"disable_direct_load,omitempty"`
|
||||
// ScopeTagIDs records the logical tag scope selected by the user. For
|
||||
// document KBs this is kept for tracing after the relation-table lookup has
|
||||
// been resolved to KnowledgeIDs; TagIDs remains the physical index filter.
|
||||
ScopeTagIDs []string `json:"scope_tag_ids,omitempty"`
|
||||
// DisableRecallThresholds keeps recall broad inside an already constrained,
|
||||
// user-selected scope. The reranker still orders candidates, but vector and
|
||||
// keyword thresholds cannot erase the whole explicit scope before reranking.
|
||||
DisableRecallThresholds bool `json:"disable_recall_thresholds,omitempty"`
|
||||
}
|
||||
|
||||
// SearchTargets is a list of search targets, pre-computed at request entry point
|
||||
type SearchTargets []*SearchTarget
|
||||
|
||||
// RecallThresholds returns the effective recall thresholds for this target.
|
||||
func (st *SearchTarget) RecallThresholds(vectorThreshold, keywordThreshold float64) (float64, float64) {
|
||||
if st != nil && st.DisableRecallThresholds {
|
||||
return 0, 0
|
||||
}
|
||||
return vectorThreshold, keywordThreshold
|
||||
}
|
||||
|
||||
// HasRecallThresholdOverride reports whether any target represents an
|
||||
// authoritative scope whose candidates must reach reranking before filtering.
|
||||
func (st SearchTargets) HasRecallThresholdOverride() bool {
|
||||
for _, target := range st {
|
||||
if target != nil && target.DisableRecallThresholds {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// HasKnowledgeRetrievalScope reports whether a request has any effective
|
||||
// knowledge retrieval scope. SearchTargets are the unified runtime form and
|
||||
// must be considered alongside the legacy/raw KB and knowledge ID fields so
|
||||
// tag-only mentions are not mistaken for pure chat.
|
||||
func HasKnowledgeRetrievalScope(
|
||||
searchTargets SearchTargets,
|
||||
knowledgeBaseIDs []string,
|
||||
knowledgeIDs []string,
|
||||
) bool {
|
||||
return len(searchTargets) > 0 || len(knowledgeBaseIDs) > 0 || len(knowledgeIDs) > 0
|
||||
}
|
||||
|
||||
// GetAllKnowledgeBaseIDs returns all unique knowledge base IDs from the search targets
|
||||
func (st SearchTargets) GetAllKnowledgeBaseIDs() []string {
|
||||
seen := make(map[string]bool)
|
||||
@@ -161,6 +195,7 @@ type SearchParams struct {
|
||||
DisableVectorMatch bool `json:"disable_vector_match"`
|
||||
KnowledgeIDs []string `json:"knowledge_ids"`
|
||||
TagIDs []string `json:"tag_ids"` // Tag IDs for filtering (used for FAQ priority filtering)
|
||||
ScopeTagIDs []string `json:"scope_tag_ids,omitempty"`
|
||||
OnlyRecommended bool `json:"only_recommended"`
|
||||
// KnowledgeBaseIDs overrides the single KB ID passed to HybridSearch,
|
||||
// allowing a single retrieval call to span multiple KBs that share the
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
package types
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestHasKnowledgeRetrievalScope(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
searchTargets SearchTargets
|
||||
knowledgeBaseIDs []string
|
||||
knowledgeIDs []string
|
||||
want bool
|
||||
}{
|
||||
{name: "empty", want: false},
|
||||
{name: "knowledge base IDs", knowledgeBaseIDs: []string{"kb-1"}, want: true},
|
||||
{name: "knowledge IDs", knowledgeIDs: []string{"doc-1"}, want: true},
|
||||
{
|
||||
name: "tag-only search target",
|
||||
searchTargets: SearchTargets{
|
||||
{
|
||||
Type: SearchTargetTypeKnowledgeBase,
|
||||
KnowledgeBaseID: "kb-1",
|
||||
TagIDs: []string{"tag-1"},
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "resolved document tag target",
|
||||
searchTargets: SearchTargets{
|
||||
{
|
||||
Type: SearchTargetTypeKnowledge,
|
||||
KnowledgeBaseID: "kb-1",
|
||||
KnowledgeIDs: []string{"doc-1"},
|
||||
ScopeTagIDs: []string{"tag-1"},
|
||||
DisableRecallThresholds: true,
|
||||
},
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := HasKnowledgeRetrievalScope(tt.searchTargets, tt.knowledgeBaseIDs, tt.knowledgeIDs)
|
||||
if got != tt.want {
|
||||
t.Fatalf("HasKnowledgeRetrievalScope() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSearchTargetRecallThresholds(t *testing.T) {
|
||||
normal := &SearchTarget{}
|
||||
vector, keyword := normal.RecallThresholds(0.5, 0.24)
|
||||
if vector != 0.5 || keyword != 0.24 {
|
||||
t.Fatalf("normal thresholds = (%v, %v), want (0.5, 0.24)", vector, keyword)
|
||||
}
|
||||
|
||||
explicit := &SearchTarget{DisableRecallThresholds: true}
|
||||
vector, keyword = explicit.RecallThresholds(0.5, 0.24)
|
||||
if vector != 0 || keyword != 0 {
|
||||
t.Fatalf("explicit thresholds = (%v, %v), want (0, 0)", vector, keyword)
|
||||
}
|
||||
if !(SearchTargets{explicit}).HasRecallThresholdOverride() {
|
||||
t.Fatal("expected explicit target to advertise recall threshold override")
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatManageCloneCopiesSearchTargetScope(t *testing.T) {
|
||||
original := &ChatManage{
|
||||
PipelineRequest: PipelineRequest{
|
||||
SearchTargets: SearchTargets{
|
||||
{
|
||||
Type: SearchTargetTypeKnowledge,
|
||||
KnowledgeBaseID: "kb-1",
|
||||
KnowledgeIDs: []string{"doc-1"},
|
||||
ScopeTagIDs: []string{"tag-1"},
|
||||
DisableRecallThresholds: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
cloned := original.Clone()
|
||||
if len(cloned.SearchTargets) != 1 {
|
||||
t.Fatalf("cloned search targets length = %d, want 1", len(cloned.SearchTargets))
|
||||
}
|
||||
got := cloned.SearchTargets[0]
|
||||
if !got.DisableRecallThresholds || len(got.ScopeTagIDs) != 1 || got.ScopeTagIDs[0] != "tag-1" {
|
||||
t.Fatalf("cloned target lost explicit scope: %#v", got)
|
||||
}
|
||||
|
||||
got.KnowledgeIDs[0] = "changed-doc"
|
||||
got.ScopeTagIDs[0] = "changed-tag"
|
||||
if original.SearchTargets[0].KnowledgeIDs[0] != "doc-1" || original.SearchTargets[0].ScopeTagIDs[0] != "tag-1" {
|
||||
t.Fatal("Clone() did not deep-copy search target scope slices")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user