mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-09-21 13:52:09 +08:00
feat: 支持输入框内选择知识库和文件,优化选择交互体验
This commit is contained in:
+12
-38
@@ -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
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
@@ -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>
|
||||
@@ -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',
|
||||
|
||||
@@ -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-вопросы и ответы по базе знаний',
|
||||
|
||||
@@ -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 问答",
|
||||
|
||||
@@ -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 || [];
|
||||
},
|
||||
},
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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 `;
|
||||
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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
};
|
||||
|
||||
|
||||
@@ -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
@@ -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}}
|
||||
`
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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") {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -21,6 +21,7 @@ const (
|
||||
MatchTypeRelationChunk // 关系Chunk匹配类型
|
||||
MatchTypeGraph
|
||||
MatchTypeWebSearch // 网络搜索匹配类型
|
||||
MatchTypeDirectLoad // 直接加载匹配类型
|
||||
)
|
||||
|
||||
// IndexInfo contains information about indexed content
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user