feat: 支持输入框内选择知识库和文件,优化选择交互体验

This commit is contained in:
wizardchen
2025-12-22 13:11:31 +08:00
committed by lyingbug
parent 485d8d7252
commit de1ffa7f7c
49 changed files with 2246 additions and 521 deletions
+12 -38
View File
@@ -11,30 +11,18 @@ conversation:
vector_threshold: 0.5
rerank_threshold: 0.5
rerank_top_k: 5
fallback_strategy: "fixed"
fallback_strategy: "model"
fallback_response: "抱歉,我无法回答这个问题。"
fallback_prompt: |
你是一个专业、友好的AI助手。现在用户提出的问题超出了你的知识库范围,你需要生成一个礼貌且有帮助的回复
你是一个专业、友好的AI助手。请根据你的知识直接回答用户的问题
## 回复要求
- 诚实承认你无法提供准确答案
- 简洁友好,不要过度道歉
- 可以提供相关的建议或替代方案
- 回复控制在50字以内
- 直接回答用户的问题
- 简洁清晰,言之有物
- 如果涉及实时数据或个人隐私信息,诚实说明无法获取
- 使用礼貌、专业的语气
## Few-shot示例
用户问题: 今天杭州西湖的游客数量是多少?
回复: 抱歉,我无法获取实时的杭州西湖游客数据。您可以通过杭州旅游官网或相关APP查询这一信息。
用户问题: 张教授的新论文发表了吗?
回复: 我没有张教授的最新论文信息。建议您查询学术数据库或直接联系张教授获取最新动态。
用户问题: 我的银行卡号是多少?
回复: 作为AI助手,我无法获取您的个人银行信息。请登录您的银行APP或联系银行客服获取相关信息。
## 用户当前的问题是:
## 用户的问题是:
{{.Query}}
enable_rewrite: true
enable_query_expansion: true
@@ -175,28 +163,14 @@ conversation:
## 以下是用户给出的文章相关信息:
generate_session_title_prompt: |
你是一个专业的会话标题生成助手,你的任务是为用户提问创建简洁、精准且具描述性的标题。
根据用户的问题,生成一个简短的会话标题。
## 格式要求
- 标题长度必须在10个字以内
- 标题应准确反映用户问题的核心主题
- 使用名词短语结构,避免使用问句
- 保持简洁明了,删除非必要词语
- 不要使用"关于"、"如何"等冗余词语开头
- 直接输出标题文本,不要有任何前缀、解释或标点符号
要求
- 5-10个字
- 提取核心主题
- 只输出标题,无需解释
## Few-shot示例
用户问题: 如何提高英语口语水平?
标题: 英语口语提升
用户问题: 最近上海有什么好玩的展览活动?
标题: 上海展览推荐
用户问题: 苹果手机电池不耐用怎么解决?
标题: 苹果电池优化
## 用户的问题是:
用户问题:
summary:
repeat_penalty: 1.0
temperature: 0.3
+5 -1
View File
@@ -28,7 +28,7 @@ export function useStream() {
let renderTimer: number | null = null
// 启动流式请求
const startStream = async (params: { session_id: any; query: any; knowledge_base_ids?: string[]; agent_enabled?: boolean; web_search_enabled?: boolean; summary_model_id?: string; mcp_service_ids?: string[]; method: string; url: string }) => {
const startStream = async (params: { session_id: any; query: any; knowledge_base_ids?: string[]; knowledge_ids?: string[]; agent_enabled?: boolean; web_search_enabled?: boolean; summary_model_id?: string; mcp_service_ids?: string[]; method: string; url: string }) => {
// 重置状态
output.value = '';
error.value = null;
@@ -85,6 +85,10 @@ export function useStream() {
if (params.knowledge_base_ids !== undefined && params.knowledge_base_ids.length > 0) {
postBody.knowledge_base_ids = params.knowledge_base_ids;
}
// Include knowledge_ids if provided
if (params.knowledge_ids !== undefined && params.knowledge_ids.length > 0) {
postBody.knowledge_ids = params.knowledge_ids;
}
// Include web_search_enabled if provided
if (params.web_search_enabled !== undefined) {
postBody.web_search_enabled = params.web_search_enabled;
+5 -1
View File
@@ -230,4 +230,8 @@ export interface FAQImportProgress {
export function getFAQImportProgress(taskId: string) {
return get(`/api/v1/faq/import/progress/${taskId}`);
}
}
export function searchKnowledge(keyword: string, offset = 0, limit = 20) {
return get(`/api/v1/knowledge/search?keyword=${encodeURIComponent(keyword)}&offset=${offset}&limit=${limit}`);
}
File diff suppressed because it is too large Load Diff
+243
View File
@@ -0,0 +1,243 @@
<template>
<div v-if="visible" class="mention-menu" :style="style" ref="menuRef" @click.stop @scroll="onScroll">
<!-- Knowledge Bases Group -->
<div v-if="kbItems.length > 0" class="mention-group">
<div class="mention-group-header">{{ $t('common.knowledgeBase') }}</div>
<div
v-for="(item, index) in kbItems"
:key="item.id"
class="mention-item"
:class="{ active: index === activeIndex }"
@click="$emit('select', item)"
@mouseenter="$emit('update:activeIndex', index)"
>
<div class="icon" :class="item.kbType === 'faq' ? 'faq-icon' : 'kb-icon'">
<t-icon :name="item.kbType === 'faq' ? 'chat-bubble-help' : 'folder'" />
</div>
<span class="name">{{ item.name }}</span>
<span class="count">({{ item.count || 0 }})</span>
</div>
</div>
<!-- Files Group -->
<div v-if="fileItems.length > 0" class="mention-group">
<div class="mention-group-header">{{ $t('common.file') }}</div>
<div
v-for="(item, index) in fileItems"
:key="item.id"
class="mention-item"
:class="{ active: (kbItems.length + index) === activeIndex }"
@click="$emit('select', item)"
@mouseenter="$emit('update:activeIndex', kbItems.length + index)"
>
<div class="icon file-icon">
<t-icon name="file" />
</div>
<span class="name">{{ item.name }}</span>
<span v-if="item.kbName" class="kb-name">{{ item.kbName }}</span>
</div>
<!-- Loading indicator -->
<div v-if="loading" class="loading-more">
<t-loading size="small" />
</div>
</div>
<div v-if="items.length === 0 && !loading" class="empty">
{{ $t('common.noResult') }}
</div>
</div>
</template>
<script setup lang="ts">
import { computed, watch, ref, nextTick } from 'vue';
const props = defineProps<{
visible: boolean;
style: any;
items: Array<{ id: string; name: string; type: 'kb' | 'file'; kbType?: 'document' | 'faq'; count?: number; kbName?: string }>;
activeIndex: number;
hasMore?: boolean;
loading?: boolean;
}>();
const emit = defineEmits(['select', 'update:activeIndex', 'loadMore']);
const menuRef = ref<HTMLElement | null>(null);
const kbItems = computed(() => props.items.filter(item => item.type === 'kb'));
const fileItems = computed(() => props.items.filter(item => item.type === 'file'));
const onScroll = (e: Event) => {
const target = e.target as HTMLElement;
const { scrollTop, scrollHeight, clientHeight } = target;
// Load more when scrolled to bottom (with 50px threshold)
if (scrollHeight - scrollTop - clientHeight < 50 && props.hasMore && !props.loading) {
emit('loadMore');
}
};
watch(() => props.activeIndex, (newIndex) => {
scrollToItem(newIndex);
});
watch(() => props.visible, (newVisible) => {
if (newVisible) {
nextTick(() => {
if (menuRef.value) menuRef.value.scrollTop = 0;
scrollToItem(props.activeIndex);
});
}
});
const scrollToItem = (index: number) => {
nextTick(() => {
if (!menuRef.value) return;
const items = menuRef.value.querySelectorAll('.mention-item');
if (!items || items.length <= index) return;
const activeItem = items[index] as HTMLElement;
const menu = menuRef.value;
if (activeItem) {
const menuRect = menu.getBoundingClientRect();
const itemRect = activeItem.getBoundingClientRect();
// 检查是否在上方被遮挡
if (itemRect.top < menuRect.top) {
menu.scrollTop -= (menuRect.top - itemRect.top);
}
// 检查是否在下方被遮挡
else if (itemRect.bottom > menuRect.bottom) {
menu.scrollTop += (itemRect.bottom - menuRect.bottom);
}
}
});
};
</script>
<style scoped>
.mention-menu {
position: fixed;
z-index: 10000;
background: var(--td-bg-color-container, #fff);
border: 1px solid var(--td-component-border, #e7e9eb);
border-radius: var(--td-radius-medium, 6px);
box-shadow: var(--td-shadow-2, 0 3px 14px 2px rgba(0, 0, 0, 0.05));
width: 300px;
max-height: 360px;
overflow-y: auto;
display: flex;
flex-direction: column;
padding: 4px 0;
}
.mention-group {
padding: 4px 0;
}
.mention-group:not(:last-child) {
border-bottom: 1px solid var(--td-component-border, #f0f0f0);
}
.mention-group-header {
padding: 8px 12px 4px;
font-size: var(--td-font-size-mark-small, 12px);
font-weight: 600;
color: var(--td-text-color-secondary, #999);
}
.mention-item {
display: flex;
align-items: center;
gap: 8px;
padding: 8px 12px;
margin: 0 4px;
cursor: pointer;
border-radius: var(--td-radius-default, 3px);
color: var(--td-text-color-primary, #333);
font-size: var(--td-font-size-body-medium, 14px);
font-family: var(--td-font-family, "PingFang SC");
transition: background 0.2s cubic-bezier(0.38, 0, 0.24, 1);
}
.mention-item:hover {
background: var(--td-bg-color-container-hover, #f3f3f3);
}
.mention-item.active {
background: var(--td-brand-color-light, #e9f8ec);
color: var(--td-brand-color, #07c05f);
}
.icon {
display: flex;
align-items: center;
justify-content: center;
width: 20px;
height: 20px;
border-radius: var(--td-radius-small, 2px);
flex-shrink: 0;
/* background: var(--td-bg-color-secondarycontainer, #f3f3f3); */
}
/* Document KB - Greenish */
.kb-icon {
background: rgba(16, 185, 129, 0.1);
color: #10b981;
}
/* FAQ KB - Blueish */
.faq-icon {
background: rgba(0, 82, 217, 0.1);
color: #0052d9;
}
/* File - Orange */
.file-icon {
background: rgba(237, 123, 47, 0.1);
color: #ed7b2f;
}
.mention-item.active .icon {
/* Active state keeps the colored icon but maybe adjusts background or just inherits */
background: transparent;
color: inherit;
}
.name {
flex: 1;
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
}
.count {
flex-shrink: 0;
font-size: var(--td-font-size-mark-small, 12px);
color: var(--td-text-color-secondary, #999);
}
.kb-name {
flex-shrink: 0;
max-width: 80px;
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
font-size: var(--td-font-size-mark-small, 12px);
color: var(--td-text-color-secondary, #999);
}
.empty {
padding: 24px 12px;
text-align: center;
color: var(--td-text-color-placeholder, #999);
font-size: var(--td-font-size-body-medium, 14px);
}
.loading-more {
display: flex;
justify-content: center;
padding: 8px 12px;
}
</style>
+4
View File
@@ -638,6 +638,9 @@ export default {
confirmDelete: 'Confirm Delete',
deleteSuccess: 'Deleted successfully',
deleteFailed: 'Delete failed',
file: 'File',
knowledgeBase: 'Knowledge Base',
noResult: 'No results',
},
file: {
upload: 'Upload File',
@@ -746,6 +749,7 @@ export default {
input: {
addModel: 'Add Model',
placeholder: 'Ask questions based on the knowledge base',
placeholderWithContext: 'Enter your question, will answer based on selected knowledge bases/files above',
agentMode: 'Agent Mode',
normalMode: 'Normal Mode',
normalModeDesc: 'Knowledge base RAG Q&A',
+4
View File
@@ -639,6 +639,9 @@ export default {
confirmDelete: 'Подтвердить удаление',
deleteSuccess: 'Успешно удалено',
deleteFailed: 'Ошибка удаления',
file: 'Файл',
knowledgeBase: 'База знаний',
noResult: 'Нет результатов',
},
file: {
upload: 'Загрузить файл',
@@ -1669,6 +1672,7 @@ export default {
input: {
addModel: 'Добавить модель',
placeholder: 'Задайте вопрос на основе базы знаний',
placeholderWithContext: 'Введите вопрос, ответ будет основан на выбранных выше базах знаний/файлах',
agentMode: 'Agent режим',
normalMode: 'Обычный режим',
normalModeDesc: 'RAG-вопросы и ответы по базе знаний',
+7
View File
@@ -731,6 +731,9 @@ export default {
confirmDelete: "确认删除",
deleteSuccess: "删除成功",
deleteFailed: "删除失败",
file: "文件",
knowledgeBase: "知识库",
noResult: "无结果",
},
file: {
upload: "上传文件",
@@ -1150,6 +1153,9 @@ export default {
messages: {
deleted: "已删除",
deleteFailed: "删除失败",
file: "文件",
knowledgeBase: "知识库",
noResult: "无结果",
},
features: {
knowledgeGraph: "已启用知识图谱",
@@ -1422,6 +1428,7 @@ export default {
input: {
addModel: "添加模型",
placeholder: "基于知识库提问",
placeholderWithContext: "输入问题,将基于上方选中的知识库/文件回答",
agentMode: "Agent 模式",
normalMode: "普通模式",
normalModeDesc: "基于知识库的 RAG 问答",
+28 -1
View File
@@ -8,6 +8,7 @@ interface Settings {
isAgentEnabled: boolean;
agentConfig: AgentConfig;
selectedKnowledgeBases: string[]; // 当前选中的知识库ID列表
selectedFiles: string[]; // 当前选中的文件ID列表
modelConfig: ModelConfig; // 模型配置
ollamaConfig: OllamaConfig; // Ollama配置
webSearchEnabled: boolean; // 网络搜索是否启用
@@ -71,6 +72,7 @@ const defaultSettings: Settings = {
use_custom_system_prompt: false
},
selectedKnowledgeBases: [], // 默认为空数组
selectedFiles: [], // 默认为空数组
modelConfig: {
chatModels: [],
embeddingModels: [],
@@ -273,5 +275,30 @@ export const useSettingsStore = defineStore("settings", {
this.settings.webSearchEnabled = enabled;
localStorage.setItem("WeKnora_settings", JSON.stringify(this.settings));
},
// File selection actions
addFile(fileId: string) {
if (!this.settings.selectedFiles) this.settings.selectedFiles = [];
if (!this.settings.selectedFiles.includes(fileId)) {
this.settings.selectedFiles.push(fileId);
localStorage.setItem("WeKnora_settings", JSON.stringify(this.settings));
}
},
removeFile(fileId: string) {
if (!this.settings.selectedFiles) return;
this.settings.selectedFiles = this.settings.selectedFiles.filter((id: string) => id !== fileId);
localStorage.setItem("WeKnora_settings", JSON.stringify(this.settings));
},
clearFiles() {
this.settings.selectedFiles = [];
localStorage.setItem("WeKnora_settings", JSON.stringify(this.settings));
},
getSelectedFiles(): string[] {
return this.settings.selectedFiles || [];
},
},
});
});
+58
View File
@@ -0,0 +1,58 @@
export interface CaretCoordinates {
top: number;
left: number;
height: number;
}
export function getCaretCoordinates(element: HTMLTextAreaElement, position: number): CaretCoordinates {
const div = document.createElement('div');
const style = window.getComputedStyle(element);
// Copy styles
const properties = [
'direction', 'boxSizing', 'width', 'height', 'overflowX', 'overflowY',
'borderTopWidth', 'borderRightWidth', 'borderBottomWidth', 'borderLeftWidth', 'borderStyle',
'paddingTop', 'paddingRight', 'paddingBottom', 'paddingLeft',
'fontStyle', 'fontVariant', 'fontWeight', 'fontStretch', 'fontSize', 'fontSizeAdjust', 'lineHeight', 'fontFamily',
'textAlign', 'textTransform', 'textIndent', 'textDecoration', 'letterSpacing', 'wordSpacing',
'tabSize', 'MozTabSize'
];
properties.forEach(prop => {
// @ts-ignore
div.style[prop] = style[prop];
});
div.style.position = 'absolute';
div.style.visibility = 'hidden';
div.style.whiteSpace = 'pre-wrap';
div.style.wordWrap = 'break-word';
div.style.top = '0';
div.style.left = '0';
// We append a special character to the end of the text to handle the case where the caret is at the end
const textContent = element.value.substring(0, position);
div.textContent = textContent;
const span = document.createElement('span');
// Use a zero-width space to simulate the caret position without adding visible width,
// but if it's at the end of a line, we might need something else.
// Standard trick is using a pipe or similar and measuring it.
span.textContent = '|';
div.appendChild(span);
document.body.appendChild(div);
const spanRect = span.getBoundingClientRect();
const divRect = div.getBoundingClientRect();
const coordinates = {
top: span.offsetTop + parseInt(style.borderTopWidth),
left: span.offsetLeft + parseInt(style.borderLeftWidth),
height: parseInt(style.lineHeight) || spanRect.height
};
document.body.removeChild(div);
return coordinates;
}
@@ -1074,9 +1074,28 @@ const renderMarkdown = (content: any): string => {
if (!content) return '';
// Ensure content is a string
const contentStr = typeof content === 'string' ? content : String(content || '');
let contentStr = typeof content === 'string' ? content : String(content || '');
if (!contentStr.trim()) return '';
// Handle streaming image syntax to prevent flickering
// Check if the content ends with an incomplete image markdown syntax like `![...` or `![...](...`
// This prevents the text from being rendered as plain text first and then jumping to an image
const lastImgStart = contentStr.lastIndexOf('![');
if (lastImgStart !== -1) {
// Only check the last occurrence to see if it's incomplete
// We check if the tail (from the last ![) contains a matching closing parenthesis )
// This is a heuristic: if the last ![ doesn't have a corresponding ), it's likely incomplete
const potentialImgTag = contentStr.slice(lastImgStart);
const hasClosingParen = potentialImgTag.includes(')');
const hasClosingBracket = potentialImgTag.includes(']');
// If we have ![ but missing ] or ), it's incomplete
// Note: This is a simple check. It might false positive on `![text] (note)` but that's rare in this context
if (!hasClosingBracket || !hasClosingParen) {
contentStr = contentStr.slice(0, lastImgStart);
}
}
try {
// Preprocess custom citation tags into safe HTML the sanitizer will allow
// Supported formats:
+3 -11
View File
@@ -311,18 +311,9 @@ const sendMsg = async (value, modelId = '') => {
// Get knowledge_base_ids from settings store (selected by user via KnowledgeBaseSelector)
const kbIds = useSettingsStoreInstance.settings.selectedKnowledgeBases || [];
const knowledgeIds = useSettingsStoreInstance.settings.selectedFiles || [];
// Validate knowledge_base_ids before sending (only when agent mode is enabled)
if (agentEnabled && kbIds.length === 0) {
MessagePlugin.warning(t('chat.selectKnowledgeBaseWarning'));
isReplying.value = false;
loading.value = false;
// 清空当前 assistant message ID
currentAssistantMessageId.value = '';
// Remove the user message that was just added
messagesList.pop();
return;
}
// Use agent-chat endpoint when agent is enabled, otherwise use knowledge-chat
const endpoint = agentEnabled ? '/api/v1/agent-chat' : '/api/v1/knowledge-chat';
@@ -333,6 +324,7 @@ const sendMsg = async (value, modelId = '') => {
await startStream({
session_id: session_id.value,
knowledge_base_ids: kbIds,
knowledge_ids: knowledgeIds,
agent_enabled: agentEnabled,
web_search_enabled: webSearchEnabled,
summary_model_id: modelId,
+3 -6
View File
@@ -44,12 +44,8 @@ const sendMsg = (value: string) => {
}
async function createNewSession(value: string) {
const selectedKbs = settingsStore.settings.selectedKnowledgeBases;
if (!selectedKbs || selectedKbs.length === 0) {
MessagePlugin.warning(t('createChat.messages.selectKnowledgeBase'));
return;
}
const selectedKbs = settingsStore.settings.selectedKnowledgeBases || [];
const selectedFiles = settingsStore.settings.selectedFiles || [];
// 构建 session 数据,包含 Agent 配置
const sessionData: any = {};
@@ -60,6 +56,7 @@ async function createNewSession(value: string) {
max_iterations: settingsStore.agentConfig.maxIterations,
temperature: settingsStore.agentConfig.temperature,
knowledge_bases: selectedKbs, // 所有选中的知识库
knowledge_ids: selectedFiles, // 所有选中的普通知识/文件
allowed_tools: settingsStore.agentConfig.allowedTools
};
+5
View File
@@ -29,6 +29,7 @@ type AgentEngine struct {
chatModel chat.Chat
eventBus *event.EventBus
knowledgeBasesInfo []*KnowledgeBaseInfo // Detailed knowledge base information for prompt
selectedDocs []*SelectedDocumentInfo // User-selected documents (via @ mention)
contextManager interfaces.ContextManager // Context manager for writing agent conversation to LLM context
sessionID string // Session ID for context management
systemPromptTemplate string // System prompt template (optional, uses default if empty)
@@ -50,6 +51,7 @@ func NewAgentEngine(
toolRegistry *tools.ToolRegistry,
eventBus *event.EventBus,
knowledgeBasesInfo []*KnowledgeBaseInfo,
selectedDocs []*SelectedDocumentInfo,
contextManager interfaces.ContextManager,
sessionID string,
systemPromptTemplate string,
@@ -63,6 +65,7 @@ func NewAgentEngine(
chatModel: chatModel,
eventBus: eventBus,
knowledgeBasesInfo: knowledgeBasesInfo,
selectedDocs: selectedDocs,
contextManager: contextManager,
sessionID: sessionID,
systemPromptTemplate: systemPromptTemplate,
@@ -99,6 +102,7 @@ func (e *AgentEngine) Execute(
systemPrompt := BuildSystemPrompt(
e.knowledgeBasesInfo,
e.config.WebSearchEnabled,
e.selectedDocs,
e.systemPromptTemplate,
)
logger.Debugf(ctx, "[Agent] SystemPrompt Length: %d characters", len(systemPrompt))
@@ -774,6 +778,7 @@ func (e *AgentEngine) streamFinalAnswerToEventBus(
systemPrompt := BuildSystemPrompt(
e.knowledgeBasesInfo,
e.config.WebSearchEnabled,
e.selectedDocs,
e.systemPromptTemplate,
)
+106 -6
View File
@@ -57,6 +57,16 @@ type RecentDocInfo struct {
FAQAnswers []string
}
// SelectedDocumentInfo contains summary information about a user-selected document (via @ mention)
// Only metadata is included; content will be fetched via tools when needed
type SelectedDocumentInfo struct {
KnowledgeID string // Knowledge ID
KnowledgeBaseID string // Knowledge base ID
Title string // Document title
FileName string // Original file name
FileType string // File type (pdf, docx, etc.)
}
// KnowledgeBaseInfo contains essential information about a knowledge base for agent prompt
type KnowledgeBaseInfo struct {
ID string
@@ -102,7 +112,8 @@ func formatKnowledgeBaseList(kbInfos []*KnowledgeBaseInfo) string {
}
var builder strings.Builder
builder.WriteString("\n")
builder.WriteString("\nThe following knowledge bases have been selected by the user for this conversation. ")
builder.WriteString("You should search within these knowledge bases to find relevant information.\n\n")
for i, kb := range kbInfos {
// Display knowledge base name and ID
builder.WriteString(fmt.Sprintf("%d. **%s** (knowledge_base_id: `%s`)\n", i+1, kb.Name, kb.ID))
@@ -184,6 +195,38 @@ func renderPromptPlaceholders(template string, knowledgeBases []*KnowledgeBaseIn
return result
}
// formatSelectedDocuments formats selected documents for the prompt (summary only, no content)
func formatSelectedDocuments(docs []*SelectedDocumentInfo) string {
if len(docs) == 0 {
return ""
}
var builder strings.Builder
builder.WriteString("\n### User Selected Documents (via @ mention)\n")
builder.WriteString("The user has explicitly selected the following documents. ")
builder.WriteString("**You should prioritize searching and retrieving information from these documents when answering.**\n")
builder.WriteString("Use `list_knowledge_chunks` with the provided Knowledge IDs to fetch their content.\n\n")
builder.WriteString("| # | Document Name | Type | Knowledge ID |\n")
builder.WriteString("|---|---------------|------|---------------|\n")
for i, doc := range docs {
title := doc.Title
if title == "" {
title = doc.FileName
}
fileType := doc.FileType
if fileType == "" {
fileType = "-"
}
builder.WriteString(fmt.Sprintf("| %d | %s | %s | `%s` |\n",
i+1, title, fileType, doc.KnowledgeID))
}
builder.WriteString("\n")
return builder.String()
}
// renderPromptPlaceholdersWithStatus renders placeholders including web search status
// Supported placeholders:
// - {{knowledge_bases}}
@@ -239,19 +282,72 @@ func BuildSystemPromptWithoutWeb(
return renderPromptPlaceholdersWithStatus(template, knowledgeBases, false, currentTime)
}
// BuildPureAgentSystemPrompt builds the system prompt for Pure Agent mode (no KBs)
func BuildPureAgentSystemPrompt(
webSearchEnabled bool,
systemPromptTemplate ...string,
) string {
var template string
if len(systemPromptTemplate) > 0 && systemPromptTemplate[0] != "" {
template = systemPromptTemplate[0]
} else {
template = PureAgentSystemPrompt
}
currentTime := time.Now().Format(time.RFC3339)
// Pass empty KB list
return renderPromptPlaceholdersWithStatus(template, []*KnowledgeBaseInfo{}, webSearchEnabled, currentTime)
}
// BuildProgressiveRAGSystemPrompt builds the progressive RAG system prompt based on web search status
// This is the main function to use - it automatically selects the appropriate version
func BuildSystemPrompt(
knowledgeBases []*KnowledgeBaseInfo,
webSearchEnabled bool,
selectedDocs []*SelectedDocumentInfo,
systemPromptTemplate ...string,
) string {
if webSearchEnabled {
return BuildSystemPromptWithWeb(knowledgeBases, systemPromptTemplate...)
var basePrompt string
// If no knowledge bases, use Pure Agent prompt
if len(knowledgeBases) == 0 {
basePrompt = BuildPureAgentSystemPrompt(webSearchEnabled, systemPromptTemplate...)
} else if webSearchEnabled {
basePrompt = BuildSystemPromptWithWeb(knowledgeBases, systemPromptTemplate...)
} else {
basePrompt = BuildSystemPromptWithoutWeb(knowledgeBases, systemPromptTemplate...)
}
return BuildSystemPromptWithoutWeb(knowledgeBases, systemPromptTemplate...)
// Append selected documents section if any
if len(selectedDocs) > 0 {
basePrompt += formatSelectedDocuments(selectedDocs)
}
return basePrompt
}
// PureAgentSystemPrompt is the system prompt for Pure Agent mode (no Knowledge Bases)
var PureAgentSystemPrompt = `### Role
You are WeKnora, an intelligent assistant powered by ReAct. You operate in a Pure Agent mode without attached Knowledge Bases.
### Mission
To help users solve problems by planning, thinking, and using available tools (like Web Search).
### Workflow
1. **Analyze:** Understand the user's request.
2. **Plan:** If the task is complex, use todo_write to create a plan.
3. **Execute:** Use available tools to gather information or perform actions.
4. **Synthesize:** Provide a comprehensive answer.
### Tool Guidelines
* **web_search / web_fetch:** Use these if enabled to find information from the internet.
* **todo_write:** Use for managing multi-step tasks.
* **thinking:** Use to plan and reflect.
### System Status
Current Time: {{current_time}}
Web Search: {{web_search_status}}
`
// ProgressiveRAGSystemPromptWithWeb is the progressive RAG system prompt template with web search enabled
// This version emphasizes hybrid retrieval strategy: KB-first with web supplementation
var ProgressiveRAGSystemPromptWithWeb = `### Role
@@ -324,7 +420,9 @@ For every retrieval attempt (Phase 1 or Phase 3), follow this exact chain:
### System Status
Current Time: {{current_time}}
Knowledge Bases: {{knowledge_bases}}
### User Selected Knowledge Bases (via @ mention)
{{knowledge_bases}}
`
// ProgressiveRAGSystemPromptWithoutWeb is the progressive RAG system prompt template without web search
@@ -407,5 +505,7 @@ For every information seeking step, strictly follow this 3-step atomic unit:
### System Status
Current Time: {{current_time}}
Knowledge Bases: {{knowledge_bases}}
### User Selected Knowledge Bases (via @ mention)
{{knowledge_bases}}
`
+13 -6
View File
@@ -20,10 +20,11 @@ type GrepChunksTool struct {
db *gorm.DB
tenantID uint64
knowledgeBaseIDs []string
knowledgeIDs []string // Specific knowledge/document IDs to search within
}
// NewGrepChunksTool creates a new grep chunks tool
func NewGrepChunksTool(db *gorm.DB, tenantID uint64, knowledgeBaseIDs []string) *GrepChunksTool {
func NewGrepChunksTool(db *gorm.DB, tenantID uint64, knowledgeBaseIDs []string, knowledgeIDs []string) *GrepChunksTool {
description := `Unix-style text pattern matching tool for knowledge base chunks.
Searches for text patterns in chunk content using strict literal text matching (fixed-string search). This tool performs exact keyword lookup, not semantic search.
@@ -68,6 +69,7 @@ grep_chunks scans enabled chunks across the specified knowledge bases and return
db: db,
tenantID: tenantID,
knowledgeBaseIDs: knowledgeBaseIDs,
knowledgeIDs: knowledgeIDs,
}
}
@@ -173,11 +175,11 @@ func (t *GrepChunksTool) Execute(ctx context.Context, args map[string]interface{
// }
// }
logger.Infof(ctx, "[Tool][GrepChunks] Patterns: %v, MaxResults: %d",
patterns, maxResults)
logger.Infof(ctx, "[Tool][GrepChunks] Patterns: %v, MaxResults: %d, KnowledgeIDs: %v",
patterns, maxResults, t.knowledgeIDs)
// Build and execute query
results, totalCount, err := t.searchChunks(ctx, patterns, kbIDs)
results, totalCount, err := t.searchChunks(ctx, patterns, kbIDs, t.knowledgeIDs)
if err != nil {
logger.Errorf(ctx, "[Tool][GrepChunks] Search failed: %v", err)
return &types.ToolResult{
@@ -269,6 +271,7 @@ func (t *GrepChunksTool) searchChunks(
ctx context.Context,
patterns []string,
kbIDs []string,
knowledgeIDs []string,
) ([]chunkWithTitle, int64, error) {
// Build base query
query := t.db.Debug().WithContext(ctx).Table("chunks").
@@ -279,8 +282,12 @@ func (t *GrepChunksTool) searchChunks(
Where("chunks.deleted_at IS NULL").
Where("knowledges.deleted_at IS NULL")
// Apply knowledge base filter
if len(kbIDs) > 0 {
// Apply knowledge IDs filter (specific documents) - takes priority over KB filter
if len(knowledgeIDs) > 0 {
query = query.Where("chunks.knowledge_id IN ?", knowledgeIDs)
logger.Infof(ctx, "[Tool][GrepChunks] Filtering by %d specific knowledge IDs", len(knowledgeIDs))
} else if len(kbIDs) > 0 {
// Apply knowledge base filter only if no specific knowledge IDs
query = query.Where("chunks.knowledge_base_id IN ?", kbIDs)
}
+77 -50
View File
@@ -32,9 +32,10 @@ type searchResultWithMeta struct {
type KnowledgeSearchTool struct {
BaseTool
knowledgeBaseService interfaces.KnowledgeBaseService
knowledgeService interfaces.KnowledgeService
chunkService interfaces.ChunkService
tenantID uint64
allowedKBs []string
searchTargets types.SearchTargets // Pre-computed unified search targets
rerankModel rerank.Reranker
chatModel chat.Chat // Optional chat model for LLM-based reranking
config *config.Config // Global config for fallback values
@@ -43,9 +44,10 @@ type KnowledgeSearchTool struct {
// NewKnowledgeSearchTool creates a new knowledge search tool
func NewKnowledgeSearchTool(
knowledgeBaseService interfaces.KnowledgeBaseService,
knowledgeService interfaces.KnowledgeService,
chunkService interfaces.ChunkService,
tenantID uint64,
allowedKBs []string,
searchTargets types.SearchTargets,
rerankModel rerank.Reranker,
chatModel chat.Chat,
cfg *config.Config,
@@ -110,9 +112,10 @@ Results represent conceptual relevance, not literal keyword overlap.
return &KnowledgeSearchTool{
BaseTool: NewBaseTool("knowledge_search", description),
knowledgeBaseService: knowledgeBaseService,
knowledgeService: knowledgeService,
chunkService: chunkService,
tenantID: tenantID,
allowedKBs: allowedKBs,
searchTargets: searchTargets,
rerankModel: rerankModel,
chatModel: chatModel,
config: cfg,
@@ -155,30 +158,46 @@ func (t *KnowledgeSearchTool) Execute(ctx context.Context, args map[string]inter
argsJSON, _ := json.MarshalIndent(args, "", " ")
logger.Debugf(ctx, "[Tool][KnowledgeSearch] Input args:\n%s", string(argsJSON))
// Determine which KBs to search
var kbIDs []string
// Determine which KBs to search - user can optionally filter to specific KBs
var userSpecifiedKBs []string
if kbIDsRaw, ok := args["knowledge_base_ids"].([]interface{}); ok && len(kbIDsRaw) > 0 {
for _, id := range kbIDsRaw {
if idStr, ok := id.(string); ok && idStr != "" {
kbIDs = append(kbIDs, idStr)
userSpecifiedKBs = append(userSpecifiedKBs, idStr)
}
}
logger.Infof(ctx, "[Tool][KnowledgeSearch] User specified %d knowledge bases: %v", len(kbIDs), kbIDs)
logger.Infof(ctx, "[Tool][KnowledgeSearch] User specified %d knowledge bases: %v", len(userSpecifiedKBs), userSpecifiedKBs)
}
// If no KBs specified, use allowed KBs
if len(kbIDs) == 0 {
kbIDs = t.allowedKBs
if len(kbIDs) == 0 {
logger.Errorf(ctx, "[Tool][KnowledgeSearch] No knowledge bases available")
return &types.ToolResult{
Success: false,
Error: "no knowledge bases specified and no allowed KBs configured",
}, fmt.Errorf("no knowledge bases available")
// Use pre-computed search targets, optionally filtered by user-specified KBs
searchTargets := t.searchTargets
if len(userSpecifiedKBs) > 0 {
// Filter search targets to only include user-specified KBs
userKBSet := make(map[string]bool)
for _, kbID := range userSpecifiedKBs {
userKBSet[kbID] = true
}
logger.Infof(ctx, "[Tool][KnowledgeSearch] Using all allowed KBs (%d): %v", len(kbIDs), kbIDs)
var filteredTargets types.SearchTargets
for _, target := range t.searchTargets {
if userKBSet[target.KnowledgeBaseID] {
filteredTargets = append(filteredTargets, target)
}
}
searchTargets = filteredTargets
}
// Validate search targets
if len(searchTargets) == 0 {
logger.Errorf(ctx, "[Tool][KnowledgeSearch] No search targets available")
return &types.ToolResult{
Success: false,
Error: "no knowledge bases specified and no search targets configured",
}, fmt.Errorf("no search targets available")
}
kbIDs := searchTargets.GetAllKnowledgeBaseIDs()
logger.Infof(ctx, "[Tool][KnowledgeSearch] Using %d search targets across %d KBs", len(searchTargets), len(kbIDs))
// Parse query parameter
var queries []string
if queriesRaw, ok := args["queries"].([]interface{}); ok && len(queriesRaw) > 0 {
@@ -256,11 +275,12 @@ func (t *KnowledgeSearchTool) Execute(ctx context.Context, args map[string]inter
minScore,
)
// Execute concurrent search (hybrid search handles both vector and keyword)
logger.Infof(ctx, "[Tool][KnowledgeSearch] Starting concurrent search across %d KBs", len(kbIDs))
// Execute concurrent search using pre-computed search targets
logger.Infof(ctx, "[Tool][KnowledgeSearch] Starting concurrent search with %d search targets",
len(searchTargets))
kbTypeMap := t.getKnowledgeBaseTypes(ctx, kbIDs)
allResults := t.concurrentSearch(ctx, queries, kbIDs,
allResults := t.concurrentSearchByTargets(ctx, queries, searchTargets,
topK, vectorThreshold, keywordThreshold, kbTypeMap)
logger.Infof(ctx, "[Tool][KnowledgeSearch] Concurrent search completed: %d raw results", len(allResults))
@@ -410,11 +430,12 @@ func (t *KnowledgeSearchTool) getKnowledgeBaseTypes(ctx context.Context, kbIDs [
return kbTypeMap
}
// concurrentSearch executes hybrid search across multiple KBs concurrently
func (t *KnowledgeSearchTool) concurrentSearch(
// concurrentSearchByTargets executes hybrid search using pre-computed search targets
// This avoids duplicate searches when a knowledge file is already covered by its KB's full search
func (t *KnowledgeSearchTool) concurrentSearchByTargets(
ctx context.Context,
queries []string,
kbsToSearch []string,
searchTargets types.SearchTargets,
topK int,
vectorThreshold, keywordThreshold float64,
kbTypeMap map[string]string,
@@ -424,24 +445,28 @@ func (t *KnowledgeSearchTool) concurrentSearch(
allResults := make([]*searchResultWithMeta, 0)
for _, query := range queries {
// Capture query in local variable to avoid closure issues
q := query
for _, kbID := range kbsToSearch {
// Capture kbID in local variable to avoid closure issues
kb := kbID
for _, target := range searchTargets {
st := target
wg.Add(1)
go func() {
defer wg.Done()
searchParams := types.SearchParams{
QueryText: q,
MatchCount: topK,
VectorThreshold: vectorThreshold,
KeywordThreshold: keywordThreshold,
}
kbResults, err := t.knowledgeBaseService.HybridSearch(ctx, kb, searchParams)
// If target has specific knowledge IDs, add them to search params
if st.Type == types.SearchTargetTypeKnowledge {
searchParams.KnowledgeIDs = st.KnowledgeIDs
}
kbResults, err := t.knowledgeBaseService.HybridSearch(ctx, st.KnowledgeBaseID, searchParams)
if err != nil {
// Log error but continue with other KBs
logger.Warnf(ctx, "[Tool][KnowledgeSearch] Failed to search knowledge base %s: %v", kb, err)
logger.Warnf(ctx, "[Tool][KnowledgeSearch] Failed to search KB %s: %v", st.KnowledgeBaseID, err)
return
}
@@ -451,9 +476,9 @@ func (t *KnowledgeSearchTool) concurrentSearch(
allResults = append(allResults, &searchResultWithMeta{
SearchResult: r,
SourceQuery: q,
QueryType: "hybrid", // Hybrid search combines both vector and keyword
KnowledgeBaseID: kb,
KnowledgeBaseType: kbTypeMap[kb],
QueryType: "hybrid",
KnowledgeBaseID: st.KnowledgeBaseID,
KnowledgeBaseType: kbTypeMap[st.KnowledgeBaseID],
})
}
mu.Unlock()
@@ -470,34 +495,36 @@ func (t *KnowledgeSearchTool) rerankResults(
query string,
results []*searchResultWithMeta,
) ([]*searchResultWithMeta, error) {
// Separate FAQ and non-FAQ results. FAQ results keep original scores.
// Separate FAQ and normal results.
// FAQ results keep original scores and bypass reranking model.
faqResults := make([]*searchResultWithMeta, 0)
nonFAQResults := make([]*searchResultWithMeta, 0, len(results))
rerankCandidates := make([]*searchResultWithMeta, 0, len(results))
for _, result := range results {
// Skip reranking for FAQ results (they are explicitly matched Q&A pairs)
if result.KnowledgeBaseType == types.KnowledgeBaseTypeFAQ {
faqResults = append(faqResults, result)
} else {
nonFAQResults = append(nonFAQResults, result)
rerankCandidates = append(rerankCandidates, result)
}
}
// If there are no non-FAQ results, return original list (already all FAQ)
if len(nonFAQResults) == 0 {
// If there are no candidates to rerank, return original list (already all FAQ)
if len(rerankCandidates) == 0 {
return results, nil
}
var (
rerankedNonFAQ []*searchResultWithMeta
err error
rerankedCandidates []*searchResultWithMeta
err error
)
// Apply reranking only to non-FAQ results
// Apply reranking only to candidates
// Try rerankModel first, fallback to chatModel if rerankModel fails or returns no results
if t.rerankModel != nil {
rerankedNonFAQ, err = t.rerankWithModel(ctx, query, nonFAQResults)
rerankedCandidates, err = t.rerankWithModel(ctx, query, rerankCandidates)
// If rerankModel fails or returns no results, fallback to chatModel
if err != nil || len(rerankedNonFAQ) == 0 {
if err != nil || len(rerankedCandidates) == 0 {
if err != nil {
logger.Warnf(ctx, "[Tool][KnowledgeSearch] Rerank model failed, falling back to chat model: %v", err)
} else {
@@ -507,18 +534,18 @@ func (t *KnowledgeSearchTool) rerankResults(
err = nil
// Try chatModel if available
if t.chatModel != nil {
rerankedNonFAQ, err = t.rerankWithLLM(ctx, query, nonFAQResults)
rerankedCandidates, err = t.rerankWithLLM(ctx, query, rerankCandidates)
} else {
// No fallback available, use original results
rerankedNonFAQ = nonFAQResults
rerankedCandidates = rerankCandidates
}
}
} else if t.chatModel != nil {
// No rerankModel, use chatModel directly
rerankedNonFAQ, err = t.rerankWithLLM(ctx, query, nonFAQResults)
rerankedCandidates, err = t.rerankWithLLM(ctx, query, rerankCandidates)
} else {
// No reranking available, use original results
rerankedNonFAQ = nonFAQResults
rerankedCandidates = rerankCandidates
}
if err != nil {
@@ -529,16 +556,16 @@ func (t *KnowledgeSearchTool) rerankResults(
logger.Debugf(ctx, "[Tool][KnowledgeSearch] Applying composite scoring")
// Store base scores before composite scoring
for _, result := range rerankedNonFAQ {
for _, result := range rerankedCandidates {
baseScore := result.Score
// Apply composite score
result.Score = t.compositeScore(result, result.Score, baseScore)
}
// Combine FAQ results (with original order) and reranked non-FAQ results
// Combine FAQ results (with original order) and reranked candidates
combined := make([]*searchResultWithMeta, 0, len(results))
combined = append(combined, faqResults...)
combined = append(combined, rerankedNonFAQ...)
combined = append(combined, rerankedCandidates...)
// Sort by score (descending) to keep consistent output order
sort.Slice(combined, func(i, j int) bool {
@@ -292,3 +292,58 @@ func (r *knowledgeRepository) CountKnowledgeByStatus(
return count, nil
}
// SearchKnowledge searches knowledge items by keyword across the tenant
// If keyword is empty, returns recent files
// Only returns documents from document-type knowledge bases (excludes FAQ)
// Returns (results, hasMore, error)
func (r *knowledgeRepository) SearchKnowledge(
ctx context.Context,
tenantID uint64,
keyword string,
offset, limit int,
) ([]*types.Knowledge, bool, error) {
// Use raw query to properly map knowledge_base_name
type KnowledgeWithKBName struct {
types.Knowledge
KnowledgeBaseName string `gorm:"column:knowledge_base_name"`
}
var results []KnowledgeWithKBName
query := r.db.WithContext(ctx).
Table("knowledges").
Select("knowledges.*, knowledge_bases.name as knowledge_base_name").
Joins("JOIN knowledge_bases ON knowledge_bases.id = knowledges.knowledge_base_id").
Where("knowledges.tenant_id = ?", tenantID).
Where("knowledge_bases.type = ?", types.KnowledgeBaseTypeDocument).
Where("knowledges.deleted_at IS NULL")
// If keyword is provided, filter by file_name or title
if keyword != "" {
query = query.Where("knowledges.file_name LIKE ? ", "%"+keyword+"%")
}
// Fetch limit+1 to check if there are more results
err := query.Order("knowledges.created_at DESC").
Offset(offset).
Limit(limit + 1).
Scan(&results).Error
if err != nil {
return nil, false, err
}
// Check if there are more results
hasMore := len(results) > limit
if hasMore {
results = results[:limit]
}
// Convert to []*types.Knowledge
knowledges := make([]*types.Knowledge, len(results))
for i, r := range results {
k := r.Knowledge
k.KnowledgeBaseName = r.KnowledgeBaseName
knowledges[i] = &k
}
return knowledges, hasMore, nil
}
@@ -38,6 +38,18 @@ func (r *knowledgeBaseRepository) GetKnowledgeBaseByID(ctx context.Context, id s
return &kb, nil
}
// GetKnowledgeBaseByIDs gets knowledge bases by multiple ids
func (r *knowledgeBaseRepository) GetKnowledgeBaseByIDs(ctx context.Context, ids []string) ([]*types.KnowledgeBase, error) {
if len(ids) == 0 {
return []*types.KnowledgeBase{}, nil
}
var kbs []*types.KnowledgeBase
if err := r.db.WithContext(ctx).Where("id IN ?", ids).Find(&kbs).Error; err != nil {
return nil, err
}
return kbs, nil
}
// ListKnowledgeBases lists all knowledge bases
func (r *knowledgeBaseRepository) ListKnowledgeBases(ctx context.Context) ([]*types.KnowledgeBase, error) {
var kbs []*types.KnowledgeBase
@@ -355,10 +355,16 @@ func (e *elasticsearchRepository) deleteByFieldList(ctx context.Context, field s
// getBaseConds Construct base Elasticsearch query conditions based on retrieval parameters
// It creates MUST conditions for required fields and MUST_NOT conditions for excluded fields
// KnowledgeBaseIDs and KnowledgeIDs use AND logic (search specific documents within knowledge bases)
// Returns a JSON string representing the query conditions
func (e *elasticsearchRepository) getBaseConds(params typesLocal.RetrieveParams) string {
// Build MUST conditions (positive filters)
must := make([]map[string]interface{}, 0)
// KnowledgeBaseIDs and KnowledgeIDs use AND logic
// - If only KnowledgeBaseIDs: search entire knowledge bases
// - If only KnowledgeIDs: search specific documents
// - If both: search specific documents within the knowledge bases (AND)
if len(params.KnowledgeBaseIDs) > 0 {
must = append(must, map[string]interface{}{
"terms": map[string]interface{}{
@@ -366,6 +372,13 @@ func (e *elasticsearchRepository) getBaseConds(params typesLocal.RetrieveParams)
},
})
}
if len(params.KnowledgeIDs) > 0 {
must = append(must, map[string]interface{}{
"terms": map[string]interface{}{
"knowledge_id.keyword": params.KnowledgeIDs,
},
})
}
// Build MUST_NOT conditions (negative filters)
mustNot := make([]map[string]interface{}, 0)
@@ -236,8 +236,14 @@ func (e *elasticsearchRepository) DeleteByKnowledgeIDList(ctx context.Context,
// getBaseConds creates the base query conditions for retrieval operations
// Returns a slice of Query objects with must and must_not conditions
// KnowledgeBaseIDs and KnowledgeIDs use AND logic (search specific documents within knowledge bases)
func (e *elasticsearchRepository) getBaseConds(params typesLocal.RetrieveParams) []types.Query {
must := []types.Query{}
// KnowledgeBaseIDs and KnowledgeIDs use AND logic
// - If only KnowledgeBaseIDs: search entire knowledge bases
// - If only KnowledgeIDs: search specific documents
// - If both: search specific documents within the knowledge bases (AND)
if len(params.KnowledgeBaseIDs) > 0 {
must = append(must, types.Query{Terms: &types.TermsQuery{
TermsQuery: map[string]types.TermsQueryField{
@@ -245,6 +251,14 @@ func (e *elasticsearchRepository) getBaseConds(params typesLocal.RetrieveParams)
},
}})
}
if len(params.KnowledgeIDs) > 0 {
must = append(must, types.Query{Terms: &types.TermsQuery{
TermsQuery: map[string]types.TermsQueryField{
"knowledge_id.keyword": params.KnowledgeIDs,
},
}})
}
mustNot := make([]types.Query, 0)
// Exclude disabled chunks (is_enabled = false)
// Note: Historical data without is_enabled field will be included (not matching must_not)
@@ -166,14 +166,25 @@ func (g *pgRepository) KeywordsRetrieve(ctx context.Context,
) ([]*types.RetrieveResult, error) {
logger.GetLogger(ctx).Infof("[Postgres] Keywords retrieval: query=%s, topK=%d", params.Query, params.TopK)
conds := make([]clause.Expression, 0)
// KnowledgeBaseIDs and KnowledgeIDs use AND logic
// - If only KnowledgeBaseIDs: search entire knowledge bases
// - If only KnowledgeIDs: search specific documents
// - If both: search specific documents within the knowledge bases (AND)
if len(params.KnowledgeBaseIDs) > 0 {
logger.GetLogger(ctx).Debugf("[Postgres] Filtering by knowledge base IDs: %v", params.KnowledgeBaseIDs)
// Use standard SQL IN clause instead of @@@ operator for better performance with B-tree index
conds = append(conds, clause.IN{
Column: "knowledge_base_id",
Values: common.ToInterfaceSlice(params.KnowledgeBaseIDs),
})
}
if len(params.KnowledgeIDs) > 0 {
logger.GetLogger(ctx).Debugf("[Postgres] Filtering by knowledge IDs: %v", params.KnowledgeIDs)
conds = append(conds, clause.IN{
Column: "knowledge_id",
Values: common.ToInterfaceSlice(params.KnowledgeIDs),
})
}
conds = append(conds, clause.Expr{
SQL: "id @@@ paradedb.match(field => 'content', value => ?, distance => 1)",
Vars: []interface{}{params.Query},
@@ -250,13 +261,15 @@ func (g *pgRepository) VectorRetrieve(ctx context.Context,
whereParts = append(whereParts, fmt.Sprintf("dimension = $%d", len(allVars)+1))
allVars = append(allVars, dimension)
// Knowledge base filter
// KnowledgeBaseIDs and KnowledgeIDs use AND logic
// - If only KnowledgeBaseIDs: search entire knowledge bases
// - If only KnowledgeIDs: search specific documents
// - If both: search specific documents within the knowledge bases (AND)
if len(params.KnowledgeBaseIDs) > 0 {
logger.GetLogger(ctx).Debugf(
"[Postgres] Filtering vector search by knowledge base IDs: %v",
params.KnowledgeBaseIDs,
)
// Build IN clause with proper placeholders
placeholders := make([]string, len(params.KnowledgeBaseIDs))
paramStart := len(allVars) + 1
for i := range params.KnowledgeBaseIDs {
@@ -266,6 +279,20 @@ func (g *pgRepository) VectorRetrieve(ctx context.Context,
whereParts = append(whereParts, fmt.Sprintf("knowledge_base_id IN (%s)",
strings.Join(placeholders, ", ")))
}
if len(params.KnowledgeIDs) > 0 {
logger.GetLogger(ctx).Debugf(
"[Postgres] Filtering vector search by knowledge IDs: %v",
params.KnowledgeIDs,
)
placeholders := make([]string, len(params.KnowledgeIDs))
paramStart := len(allVars) + 1
for i := range params.KnowledgeIDs {
placeholders[i] = fmt.Sprintf("$%d", paramStart+i)
allVars = append(allVars, params.KnowledgeIDs[i])
}
whereParts = append(whereParts, fmt.Sprintf("knowledge_id IN (%s)",
strings.Join(placeholders, ", ")))
}
// is_enabled filter
whereParts = append(whereParts, fmt.Sprintf("(is_enabled IS NULL OR is_enabled = $%d)", len(allVars)+1))
@@ -433,9 +433,16 @@ func (q *qdrantRepository) getBaseFilter(params types.RetrieveParams) *qdrant.Fi
// Only retrieve enabled chunks
must = append(must, qdrant.NewMatchBool(fieldIsEnabled, true))
// KnowledgeBaseIDs and KnowledgeIDs use AND logic
// - If only KnowledgeBaseIDs: search entire knowledge bases
// - If only KnowledgeIDs: search specific documents
// - If both: search specific documents within the knowledge bases (AND)
if len(params.KnowledgeBaseIDs) > 0 {
must = append(must, qdrant.NewMatchKeywords(fieldKnowledgeBaseID, params.KnowledgeBaseIDs...))
}
if len(params.KnowledgeIDs) > 0 {
must = append(must, qdrant.NewMatchKeywords(fieldKnowledgeID, params.KnowledgeIDs...))
}
if len(params.ExcludeKnowledgeIDs) > 0 {
mustNot = append(mustNot, qdrant.NewMatchKeywords(fieldKnowledgeID, params.ExcludeKnowledgeIDs...))
@@ -445,10 +452,12 @@ func (q *qdrantRepository) getBaseFilter(params types.RetrieveParams) *qdrant.Fi
mustNot = append(mustNot, qdrant.NewMatchKeywords(fieldChunkID, params.ExcludeChunkIDs...))
}
return &qdrant.Filter{
filter := &qdrant.Filter{
Must: must,
MustNot: mustNot,
}
return filter
}
// Retrieve dispatches the retrieval operation to the appropriate method based on retriever type
+86 -3
View File
@@ -141,6 +141,13 @@ func (s *agentService) CreateAgentEngine(
}
}
// Get selected documents information (user @ mentioned documents)
selectedDocs, err := s.getSelectedDocumentInfos(ctx, config.KnowledgeIDs)
if err != nil {
logger.Warnf(ctx, "Failed to get selected document details: %v", err)
selectedDocs = []*agent.SelectedDocumentInfo{}
}
systemPromptTemplate := ""
if config.UseCustomSystemPrompt {
systemPromptTemplate = config.ResolveSystemPrompt(config.WebSearchEnabled)
@@ -153,6 +160,7 @@ func (s *agentService) CreateAgentEngine(
toolRegistry,
eventBus,
kbInfos,
selectedDocs,
contextManager,
sessionID,
systemPromptTemplate,
@@ -173,6 +181,29 @@ func (s *agentService) registerTools(
) error {
// If no specific tools allowed, register default tools
allowedTools := tools.DefaultAllowedTools()
// Filter out knowledge base tools if no knowledge bases or knowledge IDs are configured
hasKnowledge := len(config.KnowledgeBases) > 0 || len(config.KnowledgeIDs) > 0
if !hasKnowledge {
filteredTools := make([]string, 0)
kbTools := map[string]bool{
"knowledge_search": true,
"grep_chunks": true,
"list_knowledge_chunks": true,
"query_knowledge_graph": true,
"get_document_info": true,
"database_query": true,
}
for _, toolName := range allowedTools {
if !kbTools[toolName] {
filteredTools = append(filteredTools, toolName)
}
}
allowedTools = filteredTools
logger.Infof(ctx, "Pure Agent Mode: Knowledge base tools disabled due to empty configuration")
}
// If web search is enabled, add web_search to allowedTools
if config.WebSearchEnabled {
allowedTools = append(allowedTools, "web_search")
@@ -203,16 +234,17 @@ func (s *agentService) registerTools(
registry.RegisterTool(
tools.NewKnowledgeSearchTool(
s.knowledgeBaseService,
s.knowledgeService,
s.chunkService,
tenantID,
config.KnowledgeBases,
config.SearchTargets,
rerankModel,
chatModel,
s.cfg,
))
case "grep_chunks":
registry.RegisterTool(tools.NewGrepChunksTool(s.db, tenantID, config.KnowledgeBases))
logger.Infof(ctx, "Registered grep_chunks tool for tenant: %d", tenantID)
registry.RegisterTool(tools.NewGrepChunksTool(s.db, tenantID, config.KnowledgeBases, config.KnowledgeIDs))
logger.Infof(ctx, "Registered grep_chunks tool for tenant: %d, KBs: %d, KnowledgeIDs: %d", tenantID, len(config.KnowledgeBases), len(config.KnowledgeIDs))
case "list_knowledge_chunks":
registry.RegisterTool(tools.NewListKnowledgeChunksTool(tenantID, s.knowledgeService, s.chunkService))
case "query_knowledge_graph":
@@ -372,3 +404,54 @@ func (s *agentService) getKnowledgeBaseInfos(ctx context.Context, kbIDs []string
return kbInfos, nil
}
// getSelectedDocumentInfos retrieves detailed information for user-selected documents (via @ mention)
// This loads the actual content of the documents to include in the system prompt
func (s *agentService) getSelectedDocumentInfos(ctx context.Context, knowledgeIDs []string) ([]*agent.SelectedDocumentInfo, error) {
if len(knowledgeIDs) == 0 {
return []*agent.SelectedDocumentInfo{}, nil
}
// Get tenant ID from context
tenantID := uint64(0)
if tid, ok := ctx.Value(types.TenantIDContextKey).(uint64); ok {
tenantID = tid
}
// Fetch knowledge metadata
knowledges, err := s.knowledgeService.GetKnowledgeBatch(ctx, tenantID, knowledgeIDs)
if err != nil {
return nil, fmt.Errorf("failed to get knowledge batch: %w", err)
}
// Build map for quick lookup
knowledgeMap := make(map[string]*types.Knowledge)
for _, k := range knowledges {
if k != nil {
knowledgeMap[k.ID] = k
}
}
selectedDocs := make([]*agent.SelectedDocumentInfo, 0, len(knowledgeIDs))
for _, kid := range knowledgeIDs {
k, ok := knowledgeMap[kid]
if !ok {
logger.Warnf(ctx, "Selected knowledge %s not found", kid)
continue
}
docInfo := &agent.SelectedDocumentInfo{
KnowledgeID: k.ID,
KnowledgeBaseID: k.KnowledgeBaseID,
Title: k.Title,
FileName: k.FileName,
FileType: k.FileType,
}
selectedDocs = append(selectedDocs, docInfo)
}
logger.Infof(ctx, "Loaded %d selected documents metadata for prompt", len(selectedDocs))
return selectedDocs, nil
}
@@ -22,6 +22,7 @@ type PluginExtractEntity struct {
modelService interfaces.ModelService // Model service for calling large language models
template *types.PromptTemplateStructured // Template for generating prompts
knowledgeBaseRepo interfaces.KnowledgeBaseRepository
knowledgeRepo interfaces.KnowledgeRepository
}
// NewPluginRewrite creates a new query rewriting plugin instance
@@ -30,12 +31,14 @@ func NewPluginExtractEntity(
eventManager *EventManager,
modelService interfaces.ModelService,
knowledgeBaseRepo interfaces.KnowledgeBaseRepository,
knowledgeRepo interfaces.KnowledgeRepository,
config *config.Config,
) *PluginExtractEntity {
res := &PluginExtractEntity{
modelService: modelService,
template: config.ExtractManager.ExtractEntity,
knowledgeBaseRepo: knowledgeBaseRepo,
knowledgeRepo: knowledgeRepo,
}
eventManager.Register(res)
return res
@@ -65,16 +68,68 @@ func (p *PluginExtractEntity) OnEvent(ctx context.Context,
return next()
}
kb, err := p.knowledgeBaseRepo.GetKnowledgeBaseByID(ctx, chatManage.KnowledgeBaseID)
// Collect all knowledge base IDs to query
kbIDSet := make(map[string]struct{})
for _, id := range chatManage.KnowledgeBaseIDs {
kbIDSet[id] = struct{}{}
}
// If KnowledgeIDs is specified, retrieve them and collect their knowledge base IDs
// Also build a mapping from KnowledgeID to KnowledgeBaseID
knowledgeToKBMap := make(map[string]string)
if len(chatManage.KnowledgeIDs) > 0 {
knowledges, err := p.knowledgeRepo.GetKnowledgeBatch(ctx, chatManage.TenantID, chatManage.KnowledgeIDs)
if err != nil {
logger.Errorf(ctx, "failed to get knowledges: %v", err)
return next()
}
for _, k := range knowledges {
kbIDSet[k.KnowledgeBaseID] = struct{}{}
knowledgeToKBMap[k.ID] = k.KnowledgeBaseID
}
}
// Convert set to slice
allKBIDs := make([]string, 0, len(kbIDSet))
for id := range kbIDSet {
allKBIDs = append(allKBIDs, id)
}
// Batch retrieve all knowledge bases
kbs, err := p.knowledgeBaseRepo.GetKnowledgeBaseByIDs(ctx, allKBIDs)
if err != nil {
logger.Errorf(ctx, "failed to get knowledge base: %v", err)
logger.Errorf(ctx, "failed to get knowledge bases: %v", err)
return next()
}
if kb.ExtractConfig == nil {
logger.Warnf(ctx, "failed to get extract config")
// Check if any knowledge base has ExtractConfig enabled and collect their IDs
enabledKBSet := make(map[string]struct{})
for _, kb := range kbs {
if kb.ExtractConfig != nil && kb.ExtractConfig.Enabled {
enabledKBSet[kb.ID] = struct{}{}
}
}
if len(enabledKBSet) == 0 {
logger.Debugf(ctx, "no knowledge base has extract config enabled")
return next()
}
// Save enabled knowledge base IDs for later use in search_entity
enabledKBIDs := make([]string, 0, len(enabledKBSet))
for id := range enabledKBSet {
enabledKBIDs = append(enabledKBIDs, id)
}
chatManage.EntityKBIDs = enabledKBIDs
// Filter knowledgeToKBMap to only include files from enabled knowledge bases
entityKnowledge := make(map[string]string)
for knowledgeID, kbID := range knowledgeToKBMap {
if _, ok := enabledKBSet[kbID]; ok {
entityKnowledge[knowledgeID] = kbID
}
}
chatManage.EntityKnowledge = entityKnowledge
template := &types.PromptTemplateStructured{
Description: p.template.Description,
Examples: p.template.Examples,
@@ -66,35 +66,54 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
return ErrGetRerankModel.WithError(err)
}
// Prepare passages for reranking
pipelineInfo(ctx, "Rerank", "build_passages", map[string]interface{}{
"candidate_cnt": len(chatManage.SearchResult),
})
// Prepare passages for reranking (excluding DirectLoad results)
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
}
// 合并Content和ImageInfo的文本内容
passage := getEnrichedPassage(ctx, result)
passages = append(passages, passage)
candidatesToRerank = append(candidatesToRerank, result)
}
// Single rerank call with RewriteQuery, use threshold degradation if no results
originalThreshold := chatManage.RerankThreshold
rerankResp := p.rerank(ctx, chatManage, rerankModel, chatManage.RewriteQuery, passages)
pipelineInfo(ctx, "Rerank", "build_passages", map[string]interface{}{
"total_cnt": len(chatManage.SearchResult),
"candidate_cnt": len(candidatesToRerank),
"direct_cnt": len(directLoadResults),
})
// If no results and threshold is high enough, try with lower threshold
if len(rerankResp) == 0 && originalThreshold > 0.3 {
degradedThreshold := originalThreshold * 0.7
if degradedThreshold < 0.3 {
degradedThreshold = 0.3
var rerankResp []rerank.RankResult
// Only call rerank model if there are candidates
if len(candidatesToRerank) > 0 {
// Single rerank call with RewriteQuery, use threshold degradation if no results
originalThreshold := chatManage.RerankThreshold
rerankResp = p.rerank(ctx, chatManage, rerankModel, chatManage.RewriteQuery, passages, candidatesToRerank)
// If no results and threshold is high enough, try with lower threshold
if len(rerankResp) == 0 && originalThreshold > 0.3 {
degradedThreshold := originalThreshold * 0.7
if degradedThreshold < 0.3 {
degradedThreshold = 0.3
}
pipelineInfo(ctx, "Rerank", "threshold_degrade", map[string]interface{}{
"original": originalThreshold,
"degraded": degradedThreshold,
})
chatManage.RerankThreshold = degradedThreshold
rerankResp = p.rerank(ctx, chatManage, rerankModel, chatManage.RewriteQuery, passages, candidatesToRerank)
// Restore original threshold
chatManage.RerankThreshold = originalThreshold
}
pipelineInfo(ctx, "Rerank", "threshold_degrade", map[string]interface{}{
"original": originalThreshold,
"degraded": degradedThreshold,
})
chatManage.RerankThreshold = degradedThreshold
rerankResp = p.rerank(ctx, chatManage, rerankModel, chatManage.RewriteQuery, passages)
// Restore original threshold
chatManage.RerankThreshold = originalThreshold
}
pipelineInfo(ctx, "Rerank", "model_response", map[string]interface{}{
@@ -114,9 +133,14 @@ 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))
reranked := make([]*types.SearchResult, 0, len(rerankResp)+len(directLoadResults))
// Process reranked results
for _, rr := range rerankResp {
sr := chatManage.SearchResult[rr.Index]
if rr.Index >= len(candidatesToRerank) {
continue
}
sr := candidatesToRerank[rr.Index]
base := sr.Score
sr.Metadata["base_score"] = fmt.Sprintf("%.4f", base)
modelScore := rr.RelevanceScore
@@ -130,6 +154,23 @@ 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)
// Assign high model score for direct load items
modelScore := 1.0
sr.Score = compositeScore(sr, modelScore, base)
pipelineInfo(ctx, "Rerank", "composite_calc_direct", map[string]interface{}{
"chunk_id": sr.ID,
"base_score": fmt.Sprintf("%.4f", base),
"model_score": fmt.Sprintf("%.4f", modelScore),
"final_score": fmt.Sprintf("%.4f", sr.Score),
"match_type": sr.MatchType,
})
reranked = append(reranked, sr)
}
final := applyMMR(ctx, reranked, chatManage, min(len(reranked), max(1, chatManage.RerankTopK)), 0.7)
chatManage.RerankResult = final
@@ -160,6 +201,7 @@ func (p *PluginRerank) OnEvent(ctx context.Context,
// rerank performs the actual reranking operation with given query and passages
func (p *PluginRerank) rerank(ctx context.Context,
chatManage *types.ChatManage, rerankModel rerank.Reranker, query string, passages []string,
candidates []*types.SearchResult,
) []rerank.RankResult {
pipelineInfo(ctx, "Rerank", "model_call", map[string]interface{}{
"query_variant": query,
@@ -179,21 +221,26 @@ func (p *PluginRerank) rerank(ctx context.Context,
"threshold": chatManage.RerankThreshold,
})
for i := range min(5, len(rerankResp)) {
pipelineInfo(ctx, "Rerank", "top_score", map[string]interface{}{
"rank": i + 1,
"score": rerankResp[i].RelevanceScore,
"chunk_id": chatManage.SearchResult[rerankResp[i].Index].ID,
"match_type": chatManage.SearchResult[rerankResp[i].Index].MatchType,
"chunk_type": chatManage.SearchResult[rerankResp[i].Index].ChunkType,
"content": chatManage.SearchResult[rerankResp[i].Index].Content,
})
if rerankResp[i].Index < len(candidates) {
pipelineInfo(ctx, "Rerank", "top_score", map[string]interface{}{
"rank": i + 1,
"score": rerankResp[i].RelevanceScore,
"chunk_id": candidates[rerankResp[i].Index].ID,
"match_type": candidates[rerankResp[i].Index].MatchType,
"chunk_type": candidates[rerankResp[i].Index].ChunkType,
"content": candidates[rerankResp[i].Index].Content,
})
}
}
// Filter results based on threshold with special handling for history matches
rankFilter := []rerank.RankResult{}
for _, result := range rerankResp {
if result.Index >= len(candidates) {
continue
}
th := chatManage.RerankThreshold
matchType := chatManage.SearchResult[result.Index].MatchType
matchType := candidates[result.Index].MatchType
if matchType == types.MatchTypeHistory {
th = math.Max(th-0.1, 0.5) // Lower threshold for history matches
}
@@ -19,6 +19,7 @@ import (
type PluginSearch struct {
knowledgeBaseService interfaces.KnowledgeBaseService
knowledgeService interfaces.KnowledgeService
chunkService interfaces.ChunkService
config *config.Config
webSearchService interfaces.WebSearchService
tenantService interfaces.TenantService
@@ -28,6 +29,7 @@ type PluginSearch struct {
func NewPluginSearch(eventManager *EventManager,
knowledgeBaseService interfaces.KnowledgeBaseService,
knowledgeService interfaces.KnowledgeService,
chunkService interfaces.ChunkService,
config *config.Config,
webSearchService interfaces.WebSearchService,
tenantService interfaces.TenantService,
@@ -36,6 +38,7 @@ func NewPluginSearch(eventManager *EventManager,
res := &PluginSearch{
knowledgeBaseService: knowledgeBaseService,
knowledgeService: knowledgeService,
chunkService: chunkService,
config: config,
webSearchService: webSearchService,
tenantService: tenantService,
@@ -54,35 +57,25 @@ func (p *PluginSearch) ActivationEvents() []types.EventType {
func (p *PluginSearch) OnEvent(ctx context.Context,
eventType types.EventType, chatManage *types.ChatManage, next func() *PluginError,
) *PluginError {
// Get knowledge base IDs list
knowledgeBaseIDs := chatManage.KnowledgeBaseIDs
if len(knowledgeBaseIDs) == 0 && chatManage.KnowledgeBaseID != "" {
// Fall back to single knowledge base
knowledgeBaseIDs = []string{chatManage.KnowledgeBaseID}
pipelineInfo(ctx, "Search", "fallback_kb", map[string]interface{}{
"session_id": chatManage.SessionID,
"kb_id": chatManage.KnowledgeBaseID,
})
}
if len(knowledgeBaseIDs) == 0 {
// Check if we have search targets
if len(chatManage.SearchTargets) == 0 && len(chatManage.KnowledgeBaseIDs) == 0 && len(chatManage.KnowledgeIDs) == 0 {
pipelineError(ctx, "Search", "kb_not_found", map[string]interface{}{
"session_id": chatManage.SessionID,
})
return ErrSearch.WithError(nil)
return nil
}
pipelineInfo(ctx, "Search", "input", map[string]interface{}{
"session_id": chatManage.SessionID,
"rewrite_query": chatManage.RewriteQuery,
"kb_ids": strings.Join(knowledgeBaseIDs, ","),
"tenant_id": chatManage.TenantID,
"web_enabled": chatManage.WebSearchEnabled,
"session_id": chatManage.SessionID,
"rewrite_query": chatManage.RewriteQuery,
"search_targets": len(chatManage.SearchTargets),
"tenant_id": chatManage.TenantID,
"web_enabled": chatManage.WebSearchEnabled,
})
// Run KB search and web search concurrently
pipelineInfo(ctx, "Search", "plan", map[string]interface{}{
"kb_count": len(knowledgeBaseIDs),
"search_targets": len(chatManage.SearchTargets),
"embedding_top_k": chatManage.EmbeddingTopK,
"vector_threshold": chatManage.VectorThreshold,
"keyword_threshold": chatManage.KeywordThreshold,
@@ -92,10 +85,10 @@ func (p *PluginSearch) OnEvent(ctx context.Context,
allResults := make([]*types.SearchResult, 0)
wg.Add(2)
// Goroutine 1: Knowledge base search (rewrite + processed)
// Goroutine 1: Knowledge base search using SearchTargets
go func() {
defer wg.Done()
kbResults := p.searchKnowledgeBases(ctx, knowledgeBaseIDs, chatManage)
kbResults := p.searchByTargets(ctx, chatManage)
if len(kbResults) > 0 {
mu.Lock()
allResults = append(allResults, kbResults...)
@@ -141,11 +134,11 @@ func (p *PluginSearch) OnEvent(ctx context.Context,
})
expTopK := max(chatManage.EmbeddingTopK*2, chatManage.RerankTopK*2)
expKwTh := chatManage.KeywordThreshold * 0.8
// Concurrent expansion retrieval across queries and KBs
// Concurrent expansion retrieval across queries and search targets
expResults := make([]*types.SearchResult, 0, expTopK*len(expansions))
var muExp sync.Mutex
var wgExp sync.WaitGroup
jobs := len(expansions) * len(knowledgeBaseIDs)
jobs := len(expansions) * len(chatManage.SearchTargets)
capSem := 16
if jobs < capSem {
capSem = jobs
@@ -159,9 +152,9 @@ func (p *PluginSearch) OnEvent(ctx context.Context,
"cap": capSem,
})
for _, q := range expansions {
for _, kbID := range knowledgeBaseIDs {
for _, target := range chatManage.SearchTargets {
wgExp.Add(1)
go func(q string, kbID string) {
go func(q string, t *types.SearchTarget) {
defer wgExp.Done()
sem <- struct{}{}
defer func() { <-sem }()
@@ -173,17 +166,21 @@ func (p *PluginSearch) OnEvent(ctx context.Context,
DisableVectorMatch: true,
DisableKeywordsMatch: false,
}
res, err := p.knowledgeBaseService.HybridSearch(ctx, kbID, paramsExp)
// Apply knowledge ID filter if this is a partial KB search
if t.Type == types.SearchTargetTypeKnowledge {
paramsExp.KnowledgeIDs = t.KnowledgeIDs
}
res, err := p.knowledgeBaseService.HybridSearch(ctx, t.KnowledgeBaseID, paramsExp)
if err != nil {
pipelineWarn(ctx, "Search", "expansion_error", map[string]interface{}{
"kb_id": kbID,
"kb_id": t.KnowledgeBaseID,
"error": err.Error(),
})
return
}
if len(res) > 0 {
pipelineInfo(ctx, "Search", "expansion_hits", map[string]interface{}{
"kb_id": kbID,
"kb_id": t.KnowledgeBaseID,
"query": q,
"hits": len(res),
})
@@ -191,7 +188,7 @@ func (p *PluginSearch) OnEvent(ctx context.Context,
expResults = append(expResults, res...)
muExp.Unlock()
}
}(q, kbID)
}(q, target)
}
}
wgExp.Wait()
@@ -308,46 +305,89 @@ func buildContentSignature(content string) string {
return searchutil.BuildContentSignature(content)
}
// searchKnowledgeBases performs KB searches across KB IDs using RewriteQuery only
func (p *PluginSearch) searchKnowledgeBases(
// searchByTargets performs KB searches using pre-computed SearchTargets
// This is the main search method that uses the unified search targets
func (p *PluginSearch) searchByTargets(
ctx context.Context,
knowledgeBaseIDs []string,
chatManage *types.ChatManage,
) []*types.SearchResult {
// Build params for rewrite query
baseParams := types.SearchParams{
QueryText: strings.TrimSpace(chatManage.RewriteQuery),
VectorThreshold: chatManage.VectorThreshold,
KeywordThreshold: chatManage.KeywordThreshold,
MatchCount: chatManage.EmbeddingTopK,
if len(chatManage.SearchTargets) == 0 {
return nil
}
var wg sync.WaitGroup
var mu sync.Mutex
var results []*types.SearchResult
// Search with rewrite query only (removed duplicate ProcessedQuery search)
for _, kbID := range knowledgeBaseIDs {
// Search each target concurrently
for _, target := range chatManage.SearchTargets {
wg.Add(1)
go func(knowledgeBaseID string) {
go func(t *types.SearchTarget) {
defer wg.Done()
res, err := p.knowledgeBaseService.HybridSearch(ctx, knowledgeBaseID, baseParams)
// List of knowledge IDs to perform vector search on
// Default to all IDs in the target
searchKnowledgeIDs := t.KnowledgeIDs
// Try direct loading for specific knowledge targets
if t.Type == types.SearchTargetTypeKnowledge {
directResults, skippedIDs := p.tryDirectChunkLoading(ctx, chatManage.TenantID, t.KnowledgeIDs)
if len(directResults) > 0 {
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 all files were loaded directly, we don't need to search anything
if len(skippedIDs) == 0 && len(t.KnowledgeIDs) > 0 {
return
}
// Otherwise, only search the files that were skipped (too large)
searchKnowledgeIDs = skippedIDs
}
// If no IDs left to search (and we are in Knowledge mode), we are done
if t.Type == types.SearchTargetTypeKnowledge && len(searchKnowledgeIDs) == 0 {
return
}
// Build params for rewrite query
params := types.SearchParams{
QueryText: strings.TrimSpace(chatManage.RewriteQuery),
VectorThreshold: chatManage.VectorThreshold,
KeywordThreshold: chatManage.KeywordThreshold,
MatchCount: chatManage.EmbeddingTopK,
}
// Apply knowledge ID filter if this is a partial KB search
if t.Type == types.SearchTargetTypeKnowledge {
params.KnowledgeIDs = searchKnowledgeIDs
}
res, err := p.knowledgeBaseService.HybridSearch(ctx, t.KnowledgeBaseID, params)
if err != nil {
pipelineWarn(ctx, "Search", "kb_search_error", map[string]interface{}{
"kb_id": knowledgeBaseID,
"query": baseParams.QueryText,
"error": err.Error(),
"kb_id": t.KnowledgeBaseID,
"target_type": t.Type,
"query": params.QueryText,
"error": err.Error(),
})
return
}
pipelineInfo(ctx, "Search", "kb_result", map[string]interface{}{
"kb_id": knowledgeBaseID,
"hit_count": len(res),
"kb_id": t.KnowledgeBaseID,
"target_type": t.Type,
"hit_count": len(res),
})
mu.Lock()
results = append(results, res...)
mu.Unlock()
}(kbID)
}(target)
}
wg.Wait()
@@ -358,6 +398,93 @@ func (p *PluginSearch) searchKnowledgeBases(
return results
}
// 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.GetKnowledgeBatch(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.Metadata = k.GetMetadata()
}
results = append(results, res)
}
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 || chatManage.TenantID <= 0 {
@@ -47,50 +47,81 @@ func (p *PluginSearchEntity) OnEvent(ctx context.Context,
return next()
}
// Get knowledge base IDs list
knowledgeBaseIDs := chatManage.KnowledgeBaseIDs
if len(knowledgeBaseIDs) == 0 && chatManage.KnowledgeBaseID != "" {
knowledgeBaseIDs = []string{chatManage.KnowledgeBaseID}
logger.Infof(ctx, "No KnowledgeBaseIDs provided, falling back to single KB: %s", chatManage.KnowledgeBaseID)
}
// Use EntityKBIDs (knowledge bases with ExtractConfig enabled)
knowledgeBaseIDs := chatManage.EntityKBIDs
// Use EntityKnowledge (KnowledgeID -> KnowledgeBaseID mapping for graph-enabled files)
entityKnowledge := chatManage.EntityKnowledge
if len(knowledgeBaseIDs) == 0 {
logger.Warnf(ctx, "No knowledge base IDs available for entity search")
if len(knowledgeBaseIDs) == 0 && len(entityKnowledge) == 0 {
logger.Warnf(ctx, "No knowledge base IDs or knowledge IDs with ExtractConfig enabled for entity search")
return next()
}
logger.Infof(ctx, "Searching entities across %d knowledge base(s): %v", len(knowledgeBaseIDs), knowledgeBaseIDs)
// Parallel search across multiple knowledge bases
// Parallel search across multiple knowledge bases and individual files
var wg sync.WaitGroup
var mu sync.Mutex
var allNodes []*types.GraphNode
var allRelations []*types.GraphRelation
for _, kbID := range knowledgeBaseIDs {
wg.Add(1)
go func(knowledgeBaseID string) {
defer wg.Done()
// If specific KnowledgeIDs are provided, search by individual files
if len(entityKnowledge) > 0 {
logger.Infof(ctx, "Searching entities across %d knowledge file(s)", len(entityKnowledge))
for knowledgeID, kbID := range entityKnowledge {
wg.Add(1)
go func(knowledgeBaseID, knowledgeID string) {
defer wg.Done()
graph, err := p.graphRepo.SearchNode(ctx, types.NameSpace{KnowledgeBase: knowledgeBaseID}, entity)
if err != nil {
logger.Errorf(ctx, "Failed to search entity in KB %s: %v", knowledgeBaseID, err)
return
}
graph, err := p.graphRepo.SearchNode(ctx, types.NameSpace{
KnowledgeBase: knowledgeBaseID,
Knowledge: knowledgeID,
}, entity)
if err != nil {
logger.Errorf(ctx, "Failed to search entity in Knowledge %s: %v", knowledgeID, err)
return
}
logger.Infof(
ctx,
"KB %s entity search result count: %d nodes, %d relations",
knowledgeBaseID,
len(graph.Node),
len(graph.Relation),
)
logger.Infof(
ctx,
"Knowledge %s entity search result count: %d nodes, %d relations",
knowledgeID,
len(graph.Node),
len(graph.Relation),
)
mu.Lock()
allNodes = append(allNodes, graph.Node...)
allRelations = append(allRelations, graph.Relation...)
mu.Unlock()
}(kbID)
mu.Lock()
allNodes = append(allNodes, graph.Node...)
allRelations = append(allRelations, graph.Relation...)
mu.Unlock()
}(kbID, knowledgeID)
}
} else {
// Otherwise, search by knowledge base
logger.Infof(ctx, "Searching entities across %d knowledge base(s): %v", len(knowledgeBaseIDs), knowledgeBaseIDs)
for _, kbID := range knowledgeBaseIDs {
wg.Add(1)
go func(knowledgeBaseID string) {
defer wg.Done()
graph, err := p.graphRepo.SearchNode(ctx, types.NameSpace{KnowledgeBase: knowledgeBaseID}, entity)
if err != nil {
logger.Errorf(ctx, "Failed to search entity in KB %s: %v", knowledgeBaseID, err)
return
}
logger.Infof(
ctx,
"KB %s entity search result count: %d nodes, %d relations",
knowledgeBaseID,
len(graph.Node),
len(graph.Relation),
)
mu.Lock()
allNodes = append(allNodes, graph.Node...)
allRelations = append(allRelations, graph.Relation...)
mu.Unlock()
}(kbID)
}
}
wg.Wait()
@@ -35,6 +35,7 @@ func NewPluginSearchParallel(
eventManager *EventManager,
knowledgeBaseService interfaces.KnowledgeBaseService,
knowledgeService interfaces.KnowledgeService,
chunkService interfaces.ChunkService,
config *config.Config,
webSearchService interfaces.WebSearchService,
tenantService interfaces.TenantService,
@@ -47,6 +48,7 @@ func NewPluginSearchParallel(
searchPlugin := &PluginSearch{
knowledgeBaseService: knowledgeBaseService,
knowledgeService: knowledgeService,
chunkService: chunkService,
config: config,
webSearchService: webSearchService,
tenantService: tenantService,
+6 -7
View File
@@ -266,7 +266,6 @@ func (e *EvaluationService) Evaluation(ctx context.Context,
StartTime: time.Now(),
},
Params: &types.ChatManage{
KnowledgeBaseID: knowledgeBaseID,
VectorThreshold: e.config.Conversation.VectorThreshold,
KeywordThreshold: e.config.Conversation.KeywordThreshold,
EmbeddingTopK: e.config.Conversation.EmbeddingTopK,
@@ -311,7 +310,7 @@ func (e *EvaluationService) Evaluation(ctx context.Context,
logger.Info(newCtx, "Evaluation task status set to running")
// Execute actual evaluation
if err := e.EvalDataset(newCtx, detail); err != nil {
if err := e.EvalDataset(newCtx, detail, knowledgeBaseID); err != nil {
detail.Task.Status = types.EvaluationStatueFailed
detail.Task.ErrMsg = err.Error()
logger.Errorf(newCtx, "Evaluation task failed: %v, task ID: %s", err, taskID)
@@ -329,7 +328,7 @@ func (e *EvaluationService) Evaluation(ctx context.Context,
// EvalDataset performs the actual evaluation of a dataset
// Processes each QA pair in parallel and records metrics
func (e *EvaluationService) EvalDataset(ctx context.Context, detail *types.EvaluationDetail) error {
func (e *EvaluationService) EvalDataset(ctx context.Context, detail *types.EvaluationDetail, knowledgeBaseID string) error {
logger.Info(ctx, "Start evaluating dataset")
logger.Infof(ctx, "Task ID: %s, Dataset ID: %s", detail.Task.ID, detail.Task.DatasetID)
@@ -352,7 +351,7 @@ func (e *EvaluationService) EvalDataset(ctx context.Context, detail *types.Evalu
logger.Infof(ctx, "Creating knowledge from %d passages", len(passages))
// Create knowledge base from passages
knowledge, err := e.knowledgeService.CreateKnowledgeFromPassage(ctx, detail.Params.KnowledgeBaseID, passages)
knowledge, err := e.knowledgeService.CreateKnowledgeFromPassage(ctx, knowledgeBaseID, passages)
if err != nil {
logger.Errorf(ctx, "Failed to create knowledge from passages: %v", err)
return err
@@ -366,12 +365,12 @@ func (e *EvaluationService) EvalDataset(ctx context.Context, detail *types.Evalu
logger.Errorf(ctx, "Failed to delete knowledge: %v, knowledge ID: %s", err, knowledge.ID)
}
logger.Infof(ctx, "Cleaning up resources - deleting knowledge base: %s", detail.Params.KnowledgeBaseID)
if err := e.knowledgeBaseService.DeleteKnowledgeBase(ctx, detail.Params.KnowledgeBaseID); err != nil {
logger.Infof(ctx, "Cleaning up resources - deleting knowledge base: %s", knowledgeBaseID)
if err := e.knowledgeBaseService.DeleteKnowledgeBase(ctx, knowledgeBaseID); err != nil {
logger.Errorf(
ctx,
"Failed to delete knowledge base: %v, knowledge base ID: %s",
err, detail.Params.KnowledgeBaseID,
err, knowledgeBaseID,
)
}
}()
+10
View File
@@ -5766,3 +5766,13 @@ func (s *knowledgeService) getOrCreateTagInTarget(
logger.Infof(ctx, "Created tag %s (ID: %s) in target KB %s", newTag.Name, newTag.ID, dstKnowledgeBaseID)
return newTag.ID
}
// SearchKnowledge searches knowledge items by keyword across the tenant
func (s *knowledgeService) SearchKnowledge(ctx context.Context, keyword string, offset, limit int) ([]*types.Knowledge, bool, error) {
tenantID, ok := ctx.Value(types.TenantIDContextKey).(uint64)
if !ok {
return nil, false, werrors.NewUnauthorizedError("Tenant ID not found in context")
}
return s.repo.SearchKnowledge(ctx, tenantID, keyword, offset, limit)
}
@@ -491,6 +491,7 @@ func (s *knowledgeBaseService) HybridSearch(ctx context.Context,
TopK: matchCount,
Threshold: params.VectorThreshold,
RetrieverType: types.VectorRetrieverType,
KnowledgeIDs: params.KnowledgeIDs,
}
// For FAQ knowledge base, use FAQ index
@@ -512,6 +513,7 @@ func (s *knowledgeBaseService) HybridSearch(ctx context.Context,
TopK: matchCount,
Threshold: params.KeywordThreshold,
RetrieverType: types.KeywordsRetrieverType,
KnowledgeIDs: params.KnowledgeIDs,
})
logger.Info(ctx, "Keyword retrieval parameters setup completed")
}
+158 -52
View File
@@ -43,6 +43,7 @@ type sessionService struct {
agentService interfaces.AgentService // Service for agent operations
sessionStorage llmcontext.ContextStorage // Session storage
knowledgeService interfaces.KnowledgeService // Service for knowledge operations
chunkService interfaces.ChunkService // Service for chunk operations
redisClient *redis.Client // Redis client for temp KB state
}
@@ -52,6 +53,7 @@ func NewSessionService(cfg *config.Config,
messageRepo interfaces.MessageRepository,
knowledgeBaseService interfaces.KnowledgeBaseService,
knowledgeService interfaces.KnowledgeService,
chunkService interfaces.ChunkService,
modelService interfaces.ModelService,
tenantService interfaces.TenantService,
eventManager *chatpipline.EventManager,
@@ -65,6 +67,7 @@ func NewSessionService(cfg *config.Config,
messageRepo: messageRepo,
knowledgeBaseService: knowledgeBaseService,
knowledgeService: knowledgeService,
chunkService: chunkService,
modelService: modelService,
tenantService: tenantService,
eventManager: eventManager,
@@ -394,6 +397,7 @@ func (s *sessionService) KnowledgeQA(
session *types.Session,
query string,
knowledgeBaseIDs []string,
knowledgeIDs []string,
assistantMessageID string,
summaryModelID string,
webSearchEnabled bool,
@@ -414,14 +418,13 @@ func (s *sessionService) KnowledgeQA(
logger.Infof(ctx, "No knowledge base IDs provided, using session default: %s", session.KnowledgeBaseID)
} else {
logger.Warnf(ctx, "Session has no associated knowledge base, session ID: %s", session.ID)
return errors.New("session has no knowledge base")
}
}
logger.Infof(ctx, "Using knowledge bases: %v", knowledgeBaseIDs)
// Determine chat model ID: prioritize request's summaryModelID, then Remote models
chatModelID, err := s.selectChatModelIDWithOverride(ctx, session, knowledgeBaseIDs, summaryModelID)
chatModelID, err := s.selectChatModelIDWithOverride(ctx, session, knowledgeBaseIDs, knowledgeIDs, summaryModelID)
if err != nil {
return err
}
@@ -511,20 +514,29 @@ func (s *sessionService) KnowledgeQA(
logger.Infof(ctx, "Fallback strategy not set, using default: %v", fallbackStrategy)
}
// Build unified search targets (computed once, used throughout pipeline)
searchTargets, err := s.buildSearchTargets(ctx, session.TenantID, knowledgeBaseIDs, knowledgeIDs)
if err != nil {
logger.Warnf(ctx, "Failed to build search targets: %v", err)
}
// Create chat management object with session settings
logger.Infof(
ctx,
"Creating chat manage object, knowledge base IDs: %v, chat model ID: %s",
"Creating chat manage object, knowledge base IDs: %v, knowledge IDs: %v, chat model ID: %s, search targets: %d",
knowledgeBaseIDs,
knowledgeIDs,
chatModelID,
len(searchTargets),
)
chatManage := &types.ChatManage{
Query: query,
RewriteQuery: query,
SessionID: session.ID,
MessageID: assistantMessageID, // NEW: For event emission in pipeline
KnowledgeBaseID: knowledgeBaseIDs[0], // For backward compatibility, use first KB ID
KnowledgeBaseIDs: knowledgeBaseIDs, // Multi-KB support
MessageID: assistantMessageID, // NEW: For event emission in pipeline
KnowledgeBaseIDs: knowledgeBaseIDs, // Multi-KB support
KnowledgeIDs: knowledgeIDs, // Specific knowledge (file) IDs
SearchTargets: searchTargets, // Pre-computed search targets
VectorThreshold: vectorThreshold,
KeywordThreshold: keywordThreshold,
EmbeddingTopK: embeddingTopK,
@@ -546,9 +558,22 @@ func (s *sessionService) KnowledgeQA(
EnableQueryExpansion: enableQueryExpansion,
}
// Determine pipeline based on knowledge bases availability
// If no knowledge bases are selected, use pure chat pipeline
var pipeline []types.EventType
if len(knowledgeBaseIDs) == 0 && len(knowledgeIDs) == 0 {
logger.Info(ctx, "No knowledge bases selected, using chat_stream pipeline")
pipeline = types.Pipline["chat_stream"]
// For pure chat, UserContent is the Query (since INTO_CHAT_MESSAGE is skipped)
chatManage.UserContent = query
} else {
logger.Info(ctx, "Knowledge bases selected, using rag_stream pipeline")
pipeline = types.Pipline["rag_stream"]
}
// Start knowledge QA event processing
logger.Info(ctx, "Triggering knowledge base question answering event")
err = s.KnowledgeQAByEvent(ctx, chatManage, types.Pipline["rag_stream"])
logger.Info(ctx, "Triggering question answering event")
err = s.KnowledgeQAByEvent(ctx, chatManage, pipeline)
if err != nil {
logger.ErrorWithFields(ctx, err, map[string]interface{}{
"session_id": session.ID,
@@ -592,6 +617,7 @@ func (s *sessionService) selectChatModelIDWithOverride(
ctx context.Context,
session *types.Session,
knowledgeBaseIDs []string,
knowledgeIDs []string,
summaryModelID string,
) (string, error) {
// First, check if request has summaryModelID override
@@ -612,33 +638,47 @@ func (s *sessionService) selectChatModelIDWithOverride(
}
// If no valid override, use default selection logic
return s.selectChatModelID(ctx, session, knowledgeBaseIDs)
return s.selectChatModelID(ctx, session, knowledgeBaseIDs, knowledgeIDs)
}
// selectChatModelID selects the appropriate chat model ID with priority for Remote models
// Priority order:
// 1. Session's SummaryModelID if it's a Remote model
// 2. First knowledge base with a Remote model
// 2. First knowledge base with a Remote model (from knowledgeBaseIDs or derived from knowledgeIDs)
// 3. Session's SummaryModelID (if not Remote)
// 4. First knowledge base's SummaryModelID
func (s *sessionService) selectChatModelID(
ctx context.Context,
session *types.Session,
knowledgeBaseIDs []string,
knowledgeIDs []string,
) (string, error) {
// First, check if session has a SummaryModelID and if it's a Remote model
if session.SummaryModelID != "" {
model, err := s.modelService.GetModelByID(ctx, session.SummaryModelID)
if err == nil && model != nil && model.Source == types.ModelSourceRemote {
logger.Infof(ctx, "Using session's Remote summary model: %s", session.SummaryModelID)
return session.SummaryModelID, nil
} else if err == nil && model != nil {
// Session has a model but it's not Remote, we'll check knowledge bases for Remote models
logger.Infof(ctx, "Session has summary model %s but it's not Remote, "+
"checking knowledge bases for Remote models", session.SummaryModelID)
return session.SummaryModelID, nil
}
// If no knowledge base IDs but have knowledge IDs, derive KB IDs from knowledge IDs
if len(knowledgeBaseIDs) == 0 && len(knowledgeIDs) > 0 {
tenantID := ctx.Value(types.TenantIDContextKey).(uint64)
knowledgeList, err := s.knowledgeService.GetKnowledgeBatch(ctx, tenantID, knowledgeIDs)
if err != nil {
logger.Warnf(ctx, "Failed to get knowledge batch for model selection: %v", err)
} else {
// Collect unique KB IDs from knowledge items
kbIDSet := make(map[string]bool)
for _, k := range knowledgeList {
if k != nil && k.KnowledgeBaseID != "" {
kbIDSet[k.KnowledgeBaseID] = true
}
}
for kbID := range kbIDSet {
knowledgeBaseIDs = append(knowledgeBaseIDs, kbID)
}
logger.Infof(ctx, "Derived %d knowledge base IDs from %d knowledge IDs for model selection",
len(knowledgeBaseIDs), len(knowledgeIDs))
}
}
// If no Remote model found from session, check knowledge bases for Remote models
if len(knowledgeBaseIDs) > 0 {
// Try to find a knowledge base with Remote model
@@ -683,12 +723,79 @@ func (s *sessionService) selectChatModelID(
}
}
// No knowledge bases - use session's SummaryModelID if available
if session.SummaryModelID != "" {
logger.Infof(ctx, "No knowledge bases, using session's summary model: %s", session.SummaryModelID)
return session.SummaryModelID, nil
}
logger.Error(ctx, "No chat model ID available")
return "", errors.New(
"no chat model ID available: session has no SummaryModelID and knowledge bases have no SummaryModelID",
"no chat model ID available: session has no SummaryModelID and no knowledge bases configured",
)
}
// buildSearchTargets computes the unified search targets from knowledgeBaseIDs and knowledgeIDs
// This is called once at the request entry point to avoid repeated queries later in the pipeline
// Logic:
// - For each knowledgeBaseID: create a SearchTargetTypeKnowledgeBase target
// - For each knowledgeID: find its knowledgeBaseID, if the KB is already in the list, skip (covered by full KB search)
// otherwise create a SearchTargetTypeKnowledge target grouped by KB
func (s *sessionService) buildSearchTargets(
ctx context.Context,
tenantID uint64,
knowledgeBaseIDs []string,
knowledgeIDs []string,
) (types.SearchTargets, error) {
var targets types.SearchTargets
// Track which KBs are fully searched
fullKBSet := make(map[string]bool)
for _, kbID := range knowledgeBaseIDs {
fullKBSet[kbID] = true
targets = append(targets, &types.SearchTarget{
Type: types.SearchTargetTypeKnowledgeBase,
KnowledgeBaseID: kbID,
})
}
// Process individual knowledge IDs
if len(knowledgeIDs) > 0 {
knowledgeList, err := s.knowledgeService.GetKnowledgeBatch(ctx, tenantID, knowledgeIDs)
if err != nil {
logger.Warnf(ctx, "Failed to get knowledge batch for search targets: %v", err)
return targets, nil // Return what we have, don't fail
}
// Group knowledge IDs by their KB, excluding those already covered by full KB search
kbToKnowledgeIDs := make(map[string][]string)
for _, k := range knowledgeList {
if k == nil || k.KnowledgeBaseID == "" {
continue
}
// Skip if this KB is already fully searched
if fullKBSet[k.KnowledgeBaseID] {
continue
}
kbToKnowledgeIDs[k.KnowledgeBaseID] = append(kbToKnowledgeIDs[k.KnowledgeBaseID], k.ID)
}
// Create SearchTargetTypeKnowledge targets for each KB with specific files
for kbID, kidList := range kbToKnowledgeIDs {
targets = append(targets, &types.SearchTarget{
Type: types.SearchTargetTypeKnowledge,
KnowledgeBaseID: kbID,
KnowledgeIDs: kidList,
})
}
}
logger.Infof(ctx, "Built %d search targets: %d full KB, %d partial KB",
len(targets), len(knowledgeBaseIDs), len(targets)-len(knowledgeBaseIDs))
return targets, nil
}
// KnowledgeQAByEvent processes knowledge QA through a series of events in the pipeline
func (s *sessionService) KnowledgeQAByEvent(ctx context.Context,
chatManage *types.ChatManage, eventList []types.EventType,
@@ -697,8 +804,8 @@ func (s *sessionService) KnowledgeQAByEvent(ctx context.Context,
defer span.End()
logger.Info(ctx, "Start processing knowledge base question answering through events")
logger.Infof(ctx, "Knowledge base question answering parameters, session ID: %s, knowledge base ID: %s, query: %s",
chatManage.SessionID, chatManage.KnowledgeBaseID, chatManage.Query)
logger.Infof(ctx, "Knowledge base question answering parameters, session ID: %s, query: %s",
chatManage.SessionID, chatManage.Query)
// Prepare method list for logging and tracing
methods := []string{}
@@ -769,7 +876,6 @@ func (s *sessionService) SearchKnowledge(ctx context.Context,
chatManage := &types.ChatManage{
Query: query,
RewriteQuery: query,
KnowledgeBaseID: knowledgeBaseID,
VectorThreshold: s.cfg.Conversation.VectorThreshold, // Use default configuration
KeywordThreshold: s.cfg.Conversation.KeywordThreshold, // Use default configuration
EmbeddingTopK: s.cfg.Conversation.EmbeddingTopK, // Use default configuration
@@ -800,10 +906,10 @@ func (s *sessionService) SearchKnowledge(ctx context.Context,
// Use specific event list, only including retrieval-related events, not LLM summarization
searchEvents := []types.EventType{
types.CHUNK_SEARCH, // Vector search
types.CHUNK_RERANK, // Rerank search results
types.CHUNK_MERGE, // Merge search results
types.FILTER_TOP_K, // Filter top K results
types.CHUNK_SEARCH, // Vector search
types.CHUNK_RERANK, // Rerank search results
types.CHUNK_MERGE, // Merge search results
types.FILTER_TOP_K, // Filter top K results
}
ctx, span := tracing.ContextWithSpan(ctx, "SessionService.SearchKnowledge")
@@ -895,13 +1001,14 @@ func (s *sessionService) AgentQA(
// Create runtime AgentConfig by merging session and tenant configs
// Tenant config provides the runtime parameters (MaxIterations, Temperature, Tools, Models)
// Session config provides KnowledgeBases
// Session config provides KnowledgeBases and KnowledgeIDs
agentConfig := &types.AgentConfig{
MaxIterations: tenantInfo.AgentConfig.MaxIterations,
ReflectionEnabled: tenantInfo.AgentConfig.ReflectionEnabled,
AllowedTools: tools.DefaultAllowedTools(),
Temperature: tenantInfo.AgentConfig.Temperature,
KnowledgeBases: session.AgentConfig.KnowledgeBases, // Use session's knowledge bases
KnowledgeIDs: session.AgentConfig.KnowledgeIDs, // Use session's knowledge IDs (individual documents)
WebSearchEnabled: session.AgentConfig.WebSearchEnabled, // Web search enabled from session config
}
@@ -919,43 +1026,42 @@ func (s *sessionService) AgentQA(
logger.Infof(ctx, "Merged agent config from tenant %d and session %s", tenantInfo.ID, sessionID)
// Log knowledge IDs if present
if len(agentConfig.KnowledgeIDs) > 0 {
logger.Infof(ctx, "Agent configured with %d individual knowledge ID(s): %v",
len(agentConfig.KnowledgeIDs), agentConfig.KnowledgeIDs)
}
// Determine knowledge bases for agent
// Priority: Session.AgentConfig.KnowledgeBases > Session.KnowledgeBaseID > All tenant knowledge bases
if len(agentConfig.KnowledgeBases) == 0 {
// Exception: If KnowledgeIDs are specified, don't auto-add all KBs (let buildSearchTargets handle it)
if len(agentConfig.KnowledgeBases) == 0 && len(agentConfig.KnowledgeIDs) == 0 {
if session.KnowledgeBaseID != "" {
// Use session's knowledge base as fallback
agentConfig.KnowledgeBases = []string{session.KnowledgeBaseID}
logger.Infof(ctx, "Using session's knowledge base for agent: %s", session.KnowledgeBaseID)
} else {
// Default to all knowledge bases under the tenant
logger.Infof(ctx, "No knowledge bases specified, fetching all knowledge bases for tenant")
allKBs, err := s.knowledgeBaseService.ListKnowledgeBases(ctx)
if err != nil {
logger.Errorf(ctx, "Failed to list knowledge bases for tenant: %v", err)
return fmt.Errorf("failed to list knowledge bases: %w", err)
}
if len(allKBs) == 0 {
logger.Warnf(ctx, "No knowledge bases available for agent session: %s", sessionID)
return errors.New("no knowledge bases available for agent")
}
// Extract knowledge base IDs
agentConfig.KnowledgeBases = make([]string, len(allKBs))
for i, kb := range allKBs {
if kb == nil {
continue
}
agentConfig.KnowledgeBases[i] = kb.ID
}
logger.Infof(ctx, "Agent defaulting to all %d knowledge base(s) in tenant: %v",
len(agentConfig.KnowledgeBases), agentConfig.KnowledgeBases)
// Allow running without knowledge bases (Pure Agent mode)
logger.Infof(ctx, "No knowledge bases specified for agent, running in pure agent mode")
}
} else if len(agentConfig.KnowledgeIDs) > 0 && len(agentConfig.KnowledgeBases) == 0 {
// User specified individual files but no KBs - don't auto-add all KBs
logger.Infof(ctx, "Agent configured with %d individual knowledge ID(s), no KB auto-expansion",
len(agentConfig.KnowledgeIDs))
} else {
logger.Infof(ctx, "Agent configured with %d knowledge base(s): %v",
len(agentConfig.KnowledgeBases), agentConfig.KnowledgeBases)
}
// Build search targets for agent (pre-compute once to avoid repeated queries)
searchTargets, err := s.buildSearchTargets(ctx, tenantInfo.ID, agentConfig.KnowledgeBases, agentConfig.KnowledgeIDs)
if err != nil {
logger.Warnf(ctx, "Failed to build search targets for agent: %v", err)
// Continue without search targets, the tool will handle empty targets
}
agentConfig.SearchTargets = searchTargets
logger.Infof(ctx, "Agent search targets built: %d targets", len(searchTargets))
summaryModelID := session.SummaryModelID
if summaryModelID == "" && tenantInfo.ConversationConfig != nil {
summaryModelID = tenantInfo.ConversationConfig.SummaryModelID
+1 -1
View File
@@ -1585,7 +1585,7 @@ func (h *InitializationHandler) checkRemoteModelConnection(ctx context.Context,
if strings.Contains(err.Error(), "401") || strings.Contains(err.Error(), "unauthorized") {
return false, "认证失败,请检查API Key"
} else if strings.Contains(err.Error(), "403") || strings.Contains(err.Error(), "forbidden") {
return false, "权限不足,请检查API Key权限"
return false, "权限不足,请检查API Key权限" + err.Error()
} else if strings.Contains(err.Error(), "404") || strings.Contains(err.Error(), "not found") {
return false, "API端点不存在,请检查Base URL"
} else if strings.Contains(err.Error(), "timeout") {
+34
View File
@@ -764,3 +764,37 @@ func (h *KnowledgeHandler) UpdateImageInfo(c *gin.Context) {
"message": "Knowledge chunk image updated successfully",
})
}
// SearchKnowledge godoc
// @Summary Search knowledge
// @Description Search knowledge files by keyword across all knowledge bases
// @Tags Knowledge
// @Accept json
// @Produce json
// @Param keyword query string false "Keyword to search"
// @Param offset query int false "Offset for pagination"
// @Param limit query int false "Limit for pagination (default 20)"
// @Success 200 {object} map[string]interface{} "Search results"
// @Failure 400 {object} errors.AppError "Invalid request"
// @Security Bearer
// @Router /knowledge/search [get]
func (h *KnowledgeHandler) SearchKnowledge(c *gin.Context) {
ctx := c.Request.Context()
keyword := c.Query("keyword")
offset, _ := strconv.Atoi(c.DefaultQuery("offset", "0"))
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "20"))
// Retrieve knowledge entries (empty keyword returns recent files)
knowledges, hasMore, err := h.kgService.SearchKnowledge(ctx, keyword, offset, limit)
if err != nil {
logger.ErrorWithFields(ctx, err, nil)
c.Error(errors.NewInternalServerError("Failed to search knowledge").WithDetails(err.Error()))
return
}
c.JSON(http.StatusOK, gin.H{
"success": true,
"data": knowledges,
"has_more": hasMore,
})
}
@@ -380,8 +380,11 @@ func (h *AgentStreamHandler) handleSessionTitle(ctx context.Context, evt event.E
return nil
}
// Use background context for title event since it may arrive after stream completion
bgCtx := context.Background()
// Append title event to stream
if err := h.streamManager.AppendEvent(h.ctx, h.sessionID, h.assistantMessageID, interfaces.StreamEvent{
if err := h.streamManager.AppendEvent(bgCtx, h.sessionID, h.assistantMessageID, interfaces.StreamEvent{
ID: evt.ID,
Type: types.ResponseTypeSessionTitle,
Content: data.Title,
@@ -392,7 +395,7 @@ func (h *AgentStreamHandler) handleSessionTitle(ctx context.Context, evt event.E
"title": data.Title,
},
}); err != nil {
logger.GetLogger(h.ctx).Error("Append session title event to stream failed", "error", err)
logger.GetLogger(h.ctx).Warn("Append session title event to stream failed (stream may have ended)", "error", err)
}
return nil
+18 -26
View File
@@ -142,18 +142,19 @@ func (h *Handler) KnowledgeQA(c *gin.Context) {
// Prepare knowledge base IDs
knowledgeBaseIDs := request.KnowledgeBaseIDs
if len(knowledgeBaseIDs) == 0 && session.KnowledgeBaseID != "" {
knowledgeBaseIDs = []string{session.KnowledgeBaseID}
logger.Infof(
ctx,
"No knowledge base IDs in request, using session default: %s",
secutils.SanitizeForLog(session.KnowledgeBaseID),
)
}
// if len(knowledgeBaseIDs) == 0 && session.KnowledgeBaseID != "" {
// knowledgeBaseIDs = []string{session.KnowledgeBaseID}
// logger.Infof(
// ctx,
// "No knowledge base IDs in request, using session default: %s",
// secutils.SanitizeForLog(session.KnowledgeBaseID),
// )
// }
// Use shared function to handle KnowledgeQA request
h.handleKnowledgeQARequest(ctx, c, session, secutils.SanitizeForLog(request.Query),
secutils.SanitizeForLogArray(knowledgeBaseIDs),
secutils.SanitizeForLogArray(request.KnowledgeIds),
assistantMessage, true, secutils.SanitizeForLog(request.SummaryModelID), request.WebSearchEnabled)
}
@@ -332,15 +333,6 @@ func (h *Handler) AgentQA(c *gin.Context) {
)
}
// Validate at least one knowledge base is available
if len(knowledgeBaseIDs) == 0 {
logger.Error(ctx, "No knowledge base available for delegation")
c.Error(
errors.NewBadRequestError("No knowledge base available. Please configure at least one knowledge base."),
)
return
}
logger.Infof(
ctx,
"Delegating to KnowledgeQA with knowledge bases: %s",
@@ -356,6 +348,9 @@ func (h *Handler) AgentQA(c *gin.Context) {
secutils.SanitizeForLogArray(
knowledgeBaseIDs,
),
secutils.SanitizeForLogArray(
request.KnowledgeIds,
),
assistantMessage,
false,
secutils.SanitizeForLog(request.SummaryModelID),
@@ -455,7 +450,8 @@ func (h *Handler) AgentQA(c *gin.Context) {
}()
// Handle events for SSE (blocking until connection is done)
h.handleAgentEventsForSSE(ctx, c, sessionID, assistantMessage.ID, requestID, eventBus)
// Wait for title only if session has no title (first message in session)
h.handleAgentEventsForSSE(ctx, c, sessionID, assistantMessage.ID, requestID, eventBus, session.Title == "")
}
// handleKnowledgeQARequest handles a KnowledgeQA request with the given parameters
@@ -466,6 +462,7 @@ func (h *Handler) handleKnowledgeQARequest(
session *types.Session,
query string,
knowledgeBaseIDs []string,
knowledgeIDs []string,
assistantMessage *types.Message,
generateTitle bool, // Whether to generate title if session has no title
summaryModelID string, // Optional summary model ID (overrides session default)
@@ -486,13 +483,6 @@ func (h *Handler) handleKnowledgeQARequest(
return
}
// Validate knowledge bases
if len(knowledgeBaseIDs) == 0 {
logger.Error(ctx, "No knowledge base ID available")
c.Error(errors.NewBadRequestError("At least one knowledge base ID is required"))
return
}
logger.Infof(ctx, "Using knowledge bases: %s", secutils.SanitizeForLog(fmt.Sprintf("%v", knowledgeBaseIDs)))
// Set headers for SSE
@@ -559,6 +549,7 @@ func (h *Handler) handleKnowledgeQARequest(
session,
query,
knowledgeBaseIDs,
knowledgeIDs,
assistantMessage.ID,
summaryModelID,
webSearchEnabled,
@@ -581,7 +572,8 @@ func (h *Handler) handleKnowledgeQARequest(
}()
// Handle events for SSE (blocking until connection is done)
h.handleAgentEventsForSSE(ctx, c, sessionID, assistantMessage.ID, requestID, eventBus)
// Wait for title only if session has no title (first message in session)
h.handleAgentEventsForSSE(ctx, c, sessionID, assistantMessage.ID, requestID, eventBus, session.Title == "")
}
// completeAssistantMessage marks an assistant message as complete and updates it
+50 -2
View File
@@ -294,11 +294,13 @@ func (h *Handler) StopSession(c *gin.Context) {
// handleAgentEventsForSSE handles agent events for SSE streaming using an existing handler
// The handler is already subscribed to events and AgentQA is already running
// This function polls StreamManager and pushes events to SSE, allowing graceful handling of disconnections
// waitForTitle: if true, wait for title event after completion (for new sessions without title)
func (h *Handler) handleAgentEventsForSSE(
ctx context.Context,
c *gin.Context,
sessionID, assistantMessageID, requestID string,
eventBus *event.EventBus,
waitForTitle bool,
) {
ticker := time.NewTicker(100 * time.Millisecond)
defer ticker.Stop()
@@ -329,6 +331,7 @@ func (h *Handler) handleAgentEventsForSSE(
// Send any new events
streamCompleted := false
titleReceived := false
for _, evt := range events {
// Check for stop event
if evt.Type == types.ResponseType(event.EventStop) {
@@ -366,6 +369,11 @@ func (h *Handler) handleAgentEventsForSSE(
streamCompleted = true
}
// Check for title event
if evt.Type == types.ResponseTypeSessionTitle {
titleReceived = true
}
// Check if connection is still alive before writing
if c.Request.Context().Err() != nil {
log.Info("Connection closed during event sending, stopping")
@@ -379,9 +387,49 @@ func (h *Handler) handleAgentEventsForSSE(
// Update offset
lastOffset = newOffset
// Check if stream is completed
// Check if stream is completed - wait for title event only if needed and not already received
if streamCompleted {
log.Infof("Stream completed for session=%s, message=%s", sessionID, assistantMessageID)
if waitForTitle && !titleReceived {
log.Infof("Stream completed for session=%s, message=%s, waiting for title event", sessionID, assistantMessageID)
// Wait up to 3 seconds for title event after completion
titleTimeout := time.After(3 * time.Second)
titleWaitLoop:
for {
select {
case <-titleTimeout:
log.Info("Title wait timeout, closing stream")
break titleWaitLoop
case <-c.Request.Context().Done():
log.Info("Connection closed while waiting for title")
return
default:
// Check for new events (title event)
events, newOff, err := h.streamManager.GetEvents(c.Request.Context(), sessionID, assistantMessageID, lastOffset)
if err != nil {
log.Warnf("Error getting events while waiting for title: %v", err)
break titleWaitLoop
}
if len(events) > 0 {
for _, evt := range events {
response := buildStreamResponse(evt, requestID)
c.SSEvent("message", response)
c.Writer.Flush()
// If we got the title, we can exit
if evt.Type == types.ResponseTypeSessionTitle {
log.Infof("Title event received: %s", evt.Content)
break titleWaitLoop
}
}
lastOffset = newOff
} else {
// No events, wait a bit before checking again
time.Sleep(100 * time.Millisecond)
}
}
}
} else {
log.Infof("Stream completed for session=%s, message=%s", sessionID, assistantMessageID)
}
sendCompletionEvent(c, requestID)
return
}
+1
View File
@@ -55,6 +55,7 @@ type GenerateTitleRequest struct {
type CreateKnowledgeQARequest struct {
Query string `json:"query" binding:"required"` // Query text for knowledge base search
KnowledgeBaseIDs []string `json:"knowledge_base_ids"` // Selected knowledge base ID for this request
KnowledgeIds []string `json:"knowledge_ids"` // Selected knowledge ID for this request
AgentEnabled bool `json:"agent_enabled"` // Whether agent mode is enabled for this request
WebSearchEnabled bool `json:"web_search_enabled"` // Whether web search is enabled for this request
SummaryModelID string `json:"summary_model_id"` // Optional summary model ID for this request (overrides session default)
+2
View File
@@ -166,6 +166,8 @@ func RegisterKnowledgeRoutes(r *gin.RouterGroup, handler *handler.KnowledgeHandl
k.PUT("/image/:id/:chunk_id", handler.UpdateImageInfo)
// 批量更新知识标签
k.PUT("/tags", handler.UpdateKnowledgeTagBatch)
// 搜索知识
k.GET("/search", handler.SearchKnowledge)
}
}
+3
View File
@@ -15,11 +15,13 @@ type AgentConfig struct {
AllowedTools []string `json:"allowed_tools"` // List of allowed tool names
Temperature float64 `json:"temperature"` // LLM temperature for agent
KnowledgeBases []string `json:"knowledge_bases"` // Accessible knowledge base IDs
KnowledgeIDs []string `json:"knowledge_ids"` // Accessible knowledge IDs (individual documents)
SystemPromptWebEnabled string `json:"system_prompt_web_enabled,omitempty"` // Custom prompt when web search is enabled
SystemPromptWebDisabled string `json:"system_prompt_web_disabled,omitempty"` // Custom prompt when web search is disabled
UseCustomSystemPrompt bool `json:"use_custom_system_prompt"` // Whether to use custom system prompt instead of default
WebSearchEnabled bool `json:"web_search_enabled"` // Whether web search tool is enabled
WebSearchMaxResults int `json:"web_search_max_results"` // Maximum number of web search results (default: 5)
SearchTargets SearchTargets `json:"-"` // Pre-computed unified search targets (runtime only)
}
// SessionAgentConfig represents session-level agent configuration
@@ -28,6 +30,7 @@ type SessionAgentConfig struct {
AgentModeEnabled bool `json:"agent_mode_enabled"` // Whether agent mode is enabled for this session
WebSearchEnabled bool `json:"web_search_enabled"` // Whether web search is enabled for this session
KnowledgeBases []string `json:"knowledge_bases"` // Accessible knowledge base IDs for this session
KnowledgeIDs []string `json:"knowledge_ids"` // Accessible knowledge IDs (individual documents) for this session
}
// Value implements driver.Valuer interface for AgentConfig
+23 -14
View File
@@ -8,12 +8,15 @@ type ChatManage struct {
RewriteQuery string `json:"rewrite_query,omitempty"` // Query after rewriting for better retrieval
History []*History `json:"history,omitempty"` // Chat history for context
KnowledgeBaseID string `json:"knowledge_base_id"` // ID of the knowledge base to search against (deprecated, use KnowledgeBaseIDs)
KnowledgeBaseIDs []string `json:"knowledge_base_ids"` // IDs of knowledge bases to search (multi-KB support)
VectorThreshold float64 `json:"vector_threshold"` // Minimum score threshold for vector search results
KeywordThreshold float64 `json:"keyword_threshold"` // Minimum score threshold for keyword search results
EmbeddingTopK int `json:"embedding_top_k"` // Number of top results to retrieve from embedding search
VectorDatabase string `json:"vector_database"` // Vector database type/name to use
KnowledgeBaseIDs []string `json:"knowledge_base_ids"` // IDs of knowledge bases to search (multi-KB support)
KnowledgeIDs []string `json:"knowledge_ids,omitempty"` // IDs of specific files to search (optional)
// SearchTargets is the pre-computed unified search targets
// Computed once at request entry point, used throughout the pipeline
SearchTargets SearchTargets `json:"-"`
VectorThreshold float64 `json:"vector_threshold"` // Minimum score threshold for vector search results
KeywordThreshold float64 `json:"keyword_threshold"` // Minimum score threshold for keyword search results
EmbeddingTopK int `json:"embedding_top_k"` // Number of top results to retrieve from embedding search
VectorDatabase string `json:"vector_database"` // Vector database type/name to use
RerankModelID string `json:"rerank_model_id"` // Model ID for reranking search results
RerankTopK int `json:"rerank_top_k"` // Number of top results after reranking
@@ -33,13 +36,15 @@ type ChatManage struct {
RewritePromptUser string `json:"rewrite_prompt_user"` // Custom user prompt for rewrite stage
// Internal fields for pipeline data processing
SearchResult []*SearchResult `json:"-"` // Results from search phase
RerankResult []*SearchResult `json:"-"` // Results after reranking
MergeResult []*SearchResult `json:"-"` // Final merged results after all processing
Entity []string `json:"-"` // List of identified entities
GraphResult *GraphData `json:"-"` // Graph data from search phase
UserContent string `json:"-"` // Processed user content
ChatResponse *ChatResponse `json:"-"` // Final response from chat model
SearchResult []*SearchResult `json:"-"` // Results from search phase
RerankResult []*SearchResult `json:"-"` // Results after reranking
MergeResult []*SearchResult `json:"-"` // Final merged results after all processing
Entity []string `json:"-"` // List of identified entities
EntityKBIDs []string `json:"-"` // Knowledge base IDs with ExtractConfig enabled
EntityKnowledge map[string]string `json:"-"` // KnowledgeID -> KnowledgeBaseID mapping for graph-enabled files
GraphResult *GraphData `json:"-"` // Graph data from search phase
UserContent string `json:"-"` // Processed user content
ChatResponse *ChatResponse `json:"-"` // Final response from chat model
// Event system for streaming responses
EventBus EventBusInterface `json:"-"` // EventBus for emitting streaming events
@@ -56,12 +61,16 @@ func (c *ChatManage) Clone() *ChatManage {
knowledgeBaseIDs := make([]string, len(c.KnowledgeBaseIDs))
copy(knowledgeBaseIDs, c.KnowledgeBaseIDs)
// Deep copy knowledge IDs slice
knowledgeIDs := make([]string, len(c.KnowledgeIDs))
copy(knowledgeIDs, c.KnowledgeIDs)
return &ChatManage{
Query: c.Query,
RewriteQuery: c.RewriteQuery,
SessionID: c.SessionID,
KnowledgeBaseID: c.KnowledgeBaseID,
KnowledgeBaseIDs: knowledgeBaseIDs,
KnowledgeIDs: knowledgeIDs,
VectorThreshold: c.VectorThreshold,
KeywordThreshold: c.KeywordThreshold,
EmbeddingTopK: c.EmbeddingTopK,
+1
View File
@@ -21,6 +21,7 @@ const (
MatchTypeRelationChunk // 关系Chunk匹配类型
MatchTypeGraph
MatchTypeWebSearch // 网络搜索匹配类型
MatchTypeDirectLoad // 直接加载匹配类型
)
// IndexInfo contains information about indexed content
+4
View File
@@ -119,6 +119,8 @@ type KnowledgeService interface {
SaveKBCloneProgress(ctx context.Context, progress *types.KBCloneProgress) error
// GetFAQImportProgress retrieves the progress of an FAQ import task
GetFAQImportProgress(ctx context.Context, taskID string) (*types.FAQImportProgress, error)
// SearchKnowledge searches knowledge items by keyword across the tenant.
SearchKnowledge(ctx context.Context, keyword string, offset, limit int) ([]*types.Knowledge, bool, error)
}
// KnowledgeRepository defines the interface for knowledge repositories.
@@ -156,4 +158,6 @@ type KnowledgeRepository interface {
CountKnowledgeByKnowledgeBaseID(ctx context.Context, tenantID uint64, kbID string) (int64, error)
// CountKnowledgeByStatus counts the number of knowledge items with the specified parse status.
CountKnowledgeByStatus(ctx context.Context, tenantID uint64, kbID string, parseStatuses []string) (int64, error)
// SearchKnowledge searches knowledge items by keyword across the tenant.
SearchKnowledge(ctx context.Context, tenantID uint64, keyword string, offset, limit int) ([]*types.Knowledge, bool, error)
}
@@ -111,6 +111,15 @@ type KnowledgeBaseRepository interface {
// - Possible errors such as record not existing, database errors, etc.
GetKnowledgeBaseByID(ctx context.Context, id string) (*types.KnowledgeBase, error)
// GetKnowledgeBaseByIDs queries knowledge bases by multiple IDs
// Parameters:
// - ctx: Context information
// - ids: List of knowledge base IDs
// Returns:
// - List of knowledge base objects
// - Possible errors such as database errors, etc.
GetKnowledgeBaseByIDs(ctx context.Context, ids []string) ([]*types.KnowledgeBase, error)
// ListKnowledgeBases lists all knowledge bases in the system
// Parameters:
// - ctx: Context information
+2 -1
View File
@@ -28,11 +28,12 @@ type SessionService interface {
GenerateTitleAsync(ctx context.Context, session *types.Session, userQuery string, eventBus *event.EventBus)
// KnowledgeQA performs knowledge-based question answering
// knowledgeBaseIDs: list of knowledge base IDs to search (supports multi-KB)
// knowledgeIDs: list of specific knowledge (file) IDs to search
// summaryModelID: optional summary model ID override (if empty, uses session/KB default)
// webSearchEnabled: whether to enable web search to supplement knowledge base results
// Events are emitted through eventBus (references, answer chunks, completion)
KnowledgeQA(ctx context.Context,
session *types.Session, query string, knowledgeBaseIDs []string,
session *types.Session, query string, knowledgeBaseIDs []string, knowledgeIDs []string,
assistantMessageID string, summaryModelID string, webSearchEnabled bool, eventBus *event.EventBus,
) error
// KnowledgeQAByEvent performs knowledge-based question answering by event
+2
View File
@@ -103,6 +103,8 @@ type Knowledge struct {
ErrorMessage string `json:"error_message"`
// Deletion time of the knowledge
DeletedAt gorm.DeletedAt `json:"deleted_at" gorm:"index"`
// Knowledge base name (not stored in database, populated on query)
KnowledgeBaseName string `json:"knowledge_base_name" gorm:"-"`
}
// GetMetadata returns the metadata as a map[string]string.
+2
View File
@@ -30,6 +30,8 @@ type RetrieveParams struct {
Embedding []float32
// Knowledge base IDs
KnowledgeBaseIDs []string
// Knowledge IDs
KnowledgeIDs []string
// Excluded knowledge IDs
ExcludeKnowledgeIDs []string
// Excluded chunk IDs
+45 -6
View File
@@ -5,6 +5,44 @@ import (
"encoding/json"
)
// SearchTargetType represents the type of search target
type SearchTargetType string
const (
// SearchTargetTypeKnowledgeBase - search entire knowledge base
SearchTargetTypeKnowledgeBase SearchTargetType = "knowledge_base"
// SearchTargetTypeKnowledge - search specific knowledge files within a knowledge base
SearchTargetTypeKnowledge SearchTargetType = "knowledge"
)
// SearchTarget represents a unified search target
// Either search an entire knowledge base, or specific knowledge files within a knowledge base
type SearchTarget struct {
// Type of search target
Type SearchTargetType `json:"type"`
// KnowledgeBaseID is the ID of the knowledge base to search
KnowledgeBaseID string `json:"knowledge_base_id"`
// KnowledgeIDs is the list of specific knowledge IDs to search within the knowledge base
// Only used when Type is SearchTargetTypeKnowledge
KnowledgeIDs []string `json:"knowledge_ids,omitempty"`
}
// SearchTargets is a list of search targets, pre-computed at request entry point
type SearchTargets []*SearchTarget
// GetAllKnowledgeBaseIDs returns all unique knowledge base IDs from the search targets
func (st SearchTargets) GetAllKnowledgeBaseIDs() []string {
seen := make(map[string]bool)
var result []string
for _, t := range st {
if !seen[t.KnowledgeBaseID] {
seen[t.KnowledgeBaseID] = true
result = append(result, t.KnowledgeBaseID)
}
}
return result
}
// SearchResult represents the search result
type SearchResult struct {
// ID
@@ -53,12 +91,13 @@ type SearchResult struct {
// SearchParams represents the search parameters
type SearchParams struct {
QueryText string `json:"query_text"`
VectorThreshold float64 `json:"vector_threshold"`
KeywordThreshold float64 `json:"keyword_threshold"`
MatchCount int `json:"match_count"`
DisableKeywordsMatch bool `json:"disable_keywords_match"`
DisableVectorMatch bool `json:"disable_vector_match"`
QueryText string `json:"query_text"`
VectorThreshold float64 `json:"vector_threshold"`
KeywordThreshold float64 `json:"keyword_threshold"`
MatchCount int `json:"match_count"`
DisableKeywordsMatch bool `json:"disable_keywords_match"`
DisableVectorMatch bool `json:"disable_vector_match"`
KnowledgeIDs []string `json:"knowledge_ids"`
}
// Value implements the driver.Valuer interface, used to convert SearchResult to database value