@@ -1002,7 +1446,7 @@ onBeforeRouteUpdate((to, from, next) => {
v-if="!isReplying"
@click="createSession(query)"
class="control-btn send-btn"
- :class="{ 'disabled': !query.length || selectedKbIds.length === 0 }"
+ :class="{ 'disabled': !query.length }"
>
@@ -1033,37 +1477,146 @@ const getImgSrc = (url: string) => {
transform: translateX(-400px);
}
+/* 富文本输入框容器 */
+.rich-input-container {
+ position: relative;
+ width: 800px;
+ background: var(--td-bg-color-container, #FFF);
+ border-radius: 12px;
+ border: 1px solid var(--td-component-border, #E7E7E7);
+ box-shadow: 0 6px 6px 0 rgba(0, 0, 0, 0.04), 0 12px 12px -1px rgba(0, 0, 0, 0.08);
+
+ &:focus-within {
+ border-color: var(--td-brand-color, #07C05F);
+ }
+}
+
+/* 选中的标签(输入框内顶部) */
+.selected-tags-inline {
+ display: flex;
+ flex-wrap: wrap;
+ gap: 6px;
+ padding: 12px 16px 8px;
+ border-bottom: 1px solid var(--td-component-border, #f0f0f0);
+}
+
+.inline-tag {
+ display: inline-flex;
+ align-items: center;
+ gap: 4px;
+ padding: 4px 8px;
+ border-radius: 6px;
+ font-size: 12px;
+ font-weight: 500;
+ cursor: default;
+ transition: all 0.15s;
+ background: var(--td-bg-color-secondarycontainer, #f3f3f3);
+ border: 1px solid transparent;
+ color: var(--td-text-color-primary, #333);
+
+ /* KB - Document (Greenish tint) */
+ &.kb-tag {
+ background: rgba(16, 185, 129, 0.08);
+ color: #059669;
+
+ .tag-icon {
+ color: #10b981;
+ }
+ }
+
+ /* KB - FAQ (Blueish tint) */
+ &.faq-tag {
+ background: rgba(0, 82, 217, 0.08);
+ color: #0052d9;
+
+ .tag-icon {
+ color: #0052d9;
+ }
+ }
+
+ /* File (Orange tint) */
+ &.file-tag {
+ background: rgba(237, 123, 47, 0.08);
+ color: #e65100;
+
+ .tag-icon {
+ color: #ed7b2f;
+ }
+ }
+
+ .tag-icon {
+ font-size: 14px;
+ display: flex;
+ align-items: center;
+ }
+
+ .tag-name {
+ max-width: 120px;
+ overflow: hidden;
+ text-overflow: ellipsis;
+ white-space: nowrap;
+ color: currentColor;
+ }
+
+ .tag-remove {
+ display: flex;
+ align-items: center;
+ justify-content: center;
+ width: 14px;
+ height: 14px;
+ margin-left: 2px;
+ border-radius: 50%;
+ font-size: 14px;
+ line-height: 1;
+ cursor: pointer;
+ opacity: 0.5;
+ transition: opacity 0.15s, background 0.15s;
+ color: currentColor;
+
+ &:hover {
+ opacity: 1;
+ background: rgba(0, 0, 0, 0.1);
+ }
+ }
+}
+
:deep(.t-textarea__inner) {
width: 100%;
- width: 800px;
- max-height: 250px !important;
- min-height: 112px !important;
+ max-height: 200px !important;
+ min-height: 120px !important;
resize: none;
- color: #000000e6;
+ color: var(--td-text-color-primary, #000000e6);
font-size: 16px;
font-weight: 400;
line-height: 24px;
- font-family: "PingFang SC";
- padding: 16px 12px 72px 16px; /* 增加底部padding为控制栏腾出更多空间(原52px -> 72px) */
- border-radius: 12px;
- border: 1px solid #E7E7E7;
+ font-family: var(--td-font-family, "PingFang SC");
+ padding: 12px 16px 56px 16px;
+ border-radius: 0 0 12px 12px;
+ border: none;
box-sizing: border-box;
- background: #FFF;
- box-shadow: 0 6px 6px 0 #0000000a, 0 12px 12px -1px #00000014;
+ background: transparent;
+ box-shadow: none;
&:focus {
- border: 1px solid #07C05F;
+ border: none;
+ box-shadow: none;
}
&::placeholder {
- color: #00000066;
- font-family: "PingFang SC";
+ color: var(--td-text-color-placeholder, #00000066);
+ font-family: var(--td-font-family, "PingFang SC");
font-size: 16px;
font-weight: 400;
line-height: 24px;
}
}
+/* 当没有选中标签时,textarea 样式 */
+.rich-input-container:not(:has(.selected-tags-inline)) :deep(.t-textarea__inner) {
+ border-radius: 12px;
+ padding-top: 16px;
+}
+
/* 控制栏 */
.control-bar {
position: absolute;
@@ -1074,12 +1627,12 @@ const getImgSrc = (url: string) => {
align-items: center;
justify-content: space-between;
gap: 8px;
- flex-wrap: wrap; /* 允许换行,避免内容过多时挤压 */
- max-height: 56px; /* 限制最大高度为两行 */
- z-index: 10; /* 提高z-index,确保在textarea滚动内容之上 */
- background: linear-gradient(to bottom, rgba(255, 255, 255, 0.6) 0%, rgba(255, 255, 255, 0.95) 40%, rgba(255, 255, 255, 1) 60%); /* 更强的渐变背景,从半透明逐渐变为完全不透明 */
- pointer-events: auto; /* 确保可以点击 */
- padding-top: 8px; /* 增加上边距,给渐变更多空间 */
+ flex-wrap: wrap;
+ max-height: 56px;
+ z-index: 10;
+ background: linear-gradient(to bottom, rgba(255, 255, 255, 0) 0%, var(--td-bg-color-container, #fff) 40%, var(--td-bg-color-container, #fff) 100%);
+ pointer-events: auto;
+ padding-top: 8px;
}
.control-left {
@@ -1087,9 +1640,8 @@ const getImgSrc = (url: string) => {
align-items: center;
gap: 8px;
flex: 1;
- overflow: hidden;
- flex-wrap: wrap; /* 允许内部元素换行 */
- min-width: 0; /* 允许缩小 */
+ flex-wrap: wrap;
+ min-width: 0;
}
.control-btn {
@@ -1099,14 +1651,14 @@ const getImgSrc = (url: string) => {
gap: 4px;
padding: 6px 10px;
border-radius: 6px;
- background: #f5f5f5;
+ background: var(--td-bg-color-secondarycontainer, #f5f5f5);
cursor: pointer;
transition: background 0.12s;
user-select: none;
flex-shrink: 0;
&:hover {
- background: #e6e6e6;
+ background: var(--td-bg-color-secondarycontainer-hover, #e6e6e6);
}
&.disabled {
@@ -1114,7 +1666,7 @@ const getImgSrc = (url: string) => {
cursor: not-allowed;
&:hover {
- background: #f5f5f5;
+ background: var(--td-bg-color-secondarycontainer, #f5f5f5);
}
}
}
@@ -1154,20 +1706,20 @@ const getImgSrc = (url: string) => {
}
&:not(.agent-active) {
- background: rgba(255, 255, 255, 0.8);
- border-color: #e0e0e0;
+ background: var(--td-bg-color-container, #fff);
+ border-color: var(--td-component-border, #e0e0e0);
.agent-mode-text {
- color: #666;
+ color: var(--td-text-color-secondary, #666);
}
.normal-mode-icon {
- color: #666;
+ color: var(--td-text-color-secondary, #666);
}
&:hover {
- background: rgba(255, 255, 255, 1);
- border-color: #b0b0b0;
+ background: var(--td-bg-color-container-hover, #fff);
+ border-color: var(--td-component-stroke, #b0b0b0);
}
}
}
@@ -1180,7 +1732,7 @@ const getImgSrc = (url: string) => {
.agent-mode-text {
font-size: 13px;
- color: #666;
+ color: var(--td-text-color-secondary, #666);
font-weight: 500;
white-space: nowrap;
margin: 0 4px;
@@ -1195,6 +1747,7 @@ const getImgSrc = (url: string) => {
height: 28px;
padding: 0 10px;
min-width: auto;
+ position: relative;
&.active {
background: rgba(16, 185, 129, 0.1);
@@ -1206,9 +1759,26 @@ const getImgSrc = (url: string) => {
}
}
+.kb-count {
+ position: absolute;
+ top: -4px;
+ right: -4px;
+ min-width: 16px;
+ height: 16px;
+ padding: 0 4px;
+ background: #07C05F;
+ color: white;
+ font-size: 10px;
+ font-weight: 600;
+ border-radius: 8px;
+ display: flex;
+ align-items: center;
+ justify-content: center;
+}
+
.kb-btn-text {
font-size: 13px;
- color: #666;
+ color: var(--td-text-color-secondary, #666);
font-weight: 500;
white-space: nowrap;
}
@@ -1240,14 +1810,14 @@ const getImgSrc = (url: string) => {
&:not(.active) {
.websearch-icon {
- color: #666;
+ color: var(--td-text-color-secondary, #666);
}
&:hover {
- background: #f0f0f0;
+ background: var(--td-bg-color-secondarycontainer-hover, #f0f0f0);
.websearch-icon {
- color: #333;
+ color: var(--td-text-color-primary, #333);
}
}
}
@@ -1259,7 +1829,7 @@ const getImgSrc = (url: string) => {
gap: 4px;
max-width: 220px;
font-size: 12px;
- color: #666;
+ color: var(--td-text-color-secondary, #666);
}
:global(.websearch-tooltip-disabled a) {
@@ -1288,65 +1858,6 @@ const getImgSrc = (url: string) => {
}
}
-.kb-tags {
- display: flex;
- align-items: center;
- gap: 6px;
- flex: 1;
- overflow-x: auto;
- scrollbar-width: none;
-
- &::-webkit-scrollbar {
- display: none;
- }
-}
-
-.kb-tag {
- display: flex;
- align-items: center;
- gap: 4px;
- padding: 4px 8px;
- background: rgba(16, 185, 129, 0.1);
- border: 1px solid rgba(16, 185, 129, 0.3);
- border-radius: 4px;
- font-size: 12px;
- color: #10b981;
- white-space: nowrap;
- transition: background 0.12s;
-
- &:hover {
- background: rgba(16, 185, 129, 0.15);
- }
-}
-
-.kb-tag-text {
- max-width: 100px;
- overflow: hidden;
- text-overflow: ellipsis;
-}
-
-.kb-tag-remove {
- cursor: pointer;
- font-weight: bold;
- font-size: 16px;
- line-height: 1;
- opacity: 0.7;
-
- &:hover {
- opacity: 1;
- }
-}
-
-.more-tag {
- background: rgba(0, 0, 0, 0.05);
- border-color: rgba(0, 0, 0, 0.1);
- color: #666;
-
- &:hover {
- background: rgba(0, 0, 0, 0.08);
- }
-}
-
.control-right {
display: flex;
align-items: center;
@@ -1412,7 +1923,7 @@ const getImgSrc = (url: string) => {
.model-display {
display: flex;
align-items: center;
- margin-left: auto; /* 推到最右边,但仍在 control-left 内 */
+ margin-left: auto;
flex-shrink: 0;
}
@@ -1482,10 +1993,10 @@ const getImgSrc = (url: string) => {
.model-selector-dropdown {
position: fixed !important;
z-index: 9999;
- background: #fff;
+ background: var(--td-bg-color-container, #fff);
border-radius: 10px;
- box-shadow: 0 6px 28px rgba(15, 23, 42, 0.08);
- border: 1px solid #e7e9eb;
+ box-shadow: var(--td-shadow-2, 0 6px 28px rgba(15, 23, 42, 0.08));
+ border: 1px solid var(--td-component-border, #e7e9eb);
overflow: hidden;
display: flex;
flex-direction: column;
@@ -1499,11 +2010,11 @@ const getImgSrc = (url: string) => {
align-items: center;
justify-content: space-between;
padding: 8px 10px;
- border-bottom: 1px solid #f2f4f5;
- background: #fafcfc;
+ border-bottom: 1px solid var(--td-component-border, #f2f4f5);
+ background: var(--td-bg-color-secondarycontainer, #fafcfc);
font-size: 12px;
font-weight: 600;
- color: #222;
+ color: var(--td-text-color-primary, #222);
}
.model-selector-content {
@@ -1521,9 +2032,9 @@ const getImgSrc = (url: string) => {
gap: 4px;
padding: 3px 8px;
border-radius: 6px;
- border: 1px solid #e1e5e6;
- background: #fff;
- color: #52575a;
+ border: 1px solid var(--td-component-border, #e1e5e6);
+ background: var(--td-bg-color-container, #fff);
+ color: var(--td-text-color-secondary, #52575a);
font-size: 11px;
font-weight: 500;
cursor: pointer;
@@ -1553,11 +2064,11 @@ const getImgSrc = (url: string) => {
}
&:hover {
- background: #f6f8f7;
+ background: var(--td-bg-color-container-hover, #f6f8f7);
}
&.selected {
- background: #eefdf5;
+ background: var(--td-brand-color-light, #eefdf5);
.model-option-name {
color: #10b981;
@@ -1566,7 +2077,7 @@ const getImgSrc = (url: string) => {
}
&.empty {
- color: #9aa0a6;
+ color: var(--td-text-color-disabled, #9aa0a6);
cursor: default;
text-align: center;
padding: 20px 8px;
@@ -1586,7 +2097,7 @@ const getImgSrc = (url: string) => {
.model-option-name {
font-size: 12px;
- color: #222;
+ color: var(--td-text-color-primary, #222);
flex: 1;
overflow: hidden;
text-overflow: ellipsis;
@@ -1596,7 +2107,7 @@ const getImgSrc = (url: string) => {
.model-option-desc {
font-size: 11px;
- color: #8b9196;
+ color: var(--td-text-color-secondary, #8b9196);
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
@@ -1635,10 +2146,10 @@ const getImgSrc = (url: string) => {
.agent-mode-selector-dropdown {
position: fixed !important;
z-index: 9999;
- background: #fff;
+ background: var(--td-bg-color-container, #fff);
border-radius: 10px;
- box-shadow: 0 6px 28px rgba(15, 23, 42, 0.08);
- border: 1px solid #e7e9eb;
+ box-shadow: var(--td-shadow-2, 0 6px 28px rgba(15, 23, 42, 0.08));
+ border: 1px solid var(--td-component-border, #e7e9eb);
overflow: hidden;
padding: 6px 8px;
min-width: 200px;
@@ -1661,7 +2172,7 @@ const getImgSrc = (url: string) => {
margin: 4px 6px;
&:hover:not(.disabled) {
- background: #f6f8f7;
+ background: var(--td-bg-color-container-hover, #f6f8f7);
}
&.disabled {
@@ -1674,7 +2185,7 @@ const getImgSrc = (url: string) => {
}
&.selected {
- background: #eefdf5;
+ background: var(--td-brand-color-light, #eefdf5);
.agent-mode-option-name {
color: #10b981;
@@ -1694,14 +2205,14 @@ const getImgSrc = (url: string) => {
.agent-mode-option-name {
font-size: 12px;
font-weight: 600;
- color: #222;
+ color: var(--td-text-color-primary, #222);
line-height: 1.4;
transition: color 0.12s;
}
.agent-mode-option-desc {
font-size: 11px;
- color: #8b9196;
+ color: var(--td-text-color-secondary, #8b9196);
line-height: 1.3;
}
@@ -1726,9 +2237,9 @@ const getImgSrc = (url: string) => {
.agent-mode-footer {
padding: 6px 10px;
- border-top: 1px solid #f2f4f5;
+ border-top: 1px solid var(--td-component-border, #f2f4f5);
margin-top: 2px;
- background: #fafcfc;
+ background: var(--td-bg-color-secondarycontainer, #fafcfc);
}
.agent-mode-link {
diff --git a/frontend/src/components/MentionSelector.vue b/frontend/src/components/MentionSelector.vue
new file mode 100644
index 000000000..84e6fba76
--- /dev/null
+++ b/frontend/src/components/MentionSelector.vue
@@ -0,0 +1,243 @@
+
+
+
+
+
+
+
diff --git a/frontend/src/i18n/locales/en-US.ts b/frontend/src/i18n/locales/en-US.ts
index 20a7cf902..5332e15d3 100644
--- a/frontend/src/i18n/locales/en-US.ts
+++ b/frontend/src/i18n/locales/en-US.ts
@@ -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',
diff --git a/frontend/src/i18n/locales/ru-RU.ts b/frontend/src/i18n/locales/ru-RU.ts
index 171ba1e37..c478f53c2 100644
--- a/frontend/src/i18n/locales/ru-RU.ts
+++ b/frontend/src/i18n/locales/ru-RU.ts
@@ -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-вопросы и ответы по базе знаний',
diff --git a/frontend/src/i18n/locales/zh-CN.ts b/frontend/src/i18n/locales/zh-CN.ts
index 2b574291f..633d881a2 100644
--- a/frontend/src/i18n/locales/zh-CN.ts
+++ b/frontend/src/i18n/locales/zh-CN.ts
@@ -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 问答",
diff --git a/frontend/src/stores/settings.ts b/frontend/src/stores/settings.ts
index 52c945924..4e2b1e97e 100644
--- a/frontend/src/stores/settings.ts
+++ b/frontend/src/stores/settings.ts
@@ -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 || [];
+ },
},
-});
\ No newline at end of file
+});
+
\ No newline at end of file
diff --git a/frontend/src/utils/caret.ts b/frontend/src/utils/caret.ts
new file mode 100644
index 000000000..512981f1c
--- /dev/null
+++ b/frontend/src/utils/caret.ts
@@ -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;
+}
diff --git a/frontend/src/views/chat/components/AgentStreamDisplay.vue b/frontend/src/views/chat/components/AgentStreamDisplay.vue
index 51b427d02..787d24c3d 100644
--- a/frontend/src/views/chat/components/AgentStreamDisplay.vue
+++ b/frontend/src/views/chat/components/AgentStreamDisplay.vue
@@ -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:
diff --git a/frontend/src/views/chat/index.vue b/frontend/src/views/chat/index.vue
index 399aeaa86..2e2a55d78 100644
--- a/frontend/src/views/chat/index.vue
+++ b/frontend/src/views/chat/index.vue
@@ -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,
diff --git a/frontend/src/views/creatChat/creatChat.vue b/frontend/src/views/creatChat/creatChat.vue
index 60c4f2719..e3a1f546a 100644
--- a/frontend/src/views/creatChat/creatChat.vue
+++ b/frontend/src/views/creatChat/creatChat.vue
@@ -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
};
diff --git a/internal/agent/engine.go b/internal/agent/engine.go
index e0364ade1..a13eec397 100644
--- a/internal/agent/engine.go
+++ b/internal/agent/engine.go
@@ -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,
)
diff --git a/internal/agent/prompts.go b/internal/agent/prompts.go
index ec50b389d..ea60f3930 100644
--- a/internal/agent/prompts.go
+++ b/internal/agent/prompts.go
@@ -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}}
`
diff --git a/internal/agent/tools/grep_chunks.go b/internal/agent/tools/grep_chunks.go
index 76f810a43..af5abe3ce 100644
--- a/internal/agent/tools/grep_chunks.go
+++ b/internal/agent/tools/grep_chunks.go
@@ -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)
}
diff --git a/internal/agent/tools/knowledge_search.go b/internal/agent/tools/knowledge_search.go
index e6b590f9e..79482b900 100644
--- a/internal/agent/tools/knowledge_search.go
+++ b/internal/agent/tools/knowledge_search.go
@@ -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 {
diff --git a/internal/application/repository/knowledge.go b/internal/application/repository/knowledge.go
index fb5133f15..4640c8fd1 100644
--- a/internal/application/repository/knowledge.go
+++ b/internal/application/repository/knowledge.go
@@ -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
+}
diff --git a/internal/application/repository/knowledgebase.go b/internal/application/repository/knowledgebase.go
index 36e125186..fff4be132 100644
--- a/internal/application/repository/knowledgebase.go
+++ b/internal/application/repository/knowledgebase.go
@@ -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
diff --git a/internal/application/repository/retriever/elasticsearch/v7/repository.go b/internal/application/repository/retriever/elasticsearch/v7/repository.go
index 60c6747fe..df624674b 100644
--- a/internal/application/repository/retriever/elasticsearch/v7/repository.go
+++ b/internal/application/repository/retriever/elasticsearch/v7/repository.go
@@ -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)
diff --git a/internal/application/repository/retriever/elasticsearch/v8/repository.go b/internal/application/repository/retriever/elasticsearch/v8/repository.go
index a73986d1a..adfe13dc6 100644
--- a/internal/application/repository/retriever/elasticsearch/v8/repository.go
+++ b/internal/application/repository/retriever/elasticsearch/v8/repository.go
@@ -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)
diff --git a/internal/application/repository/retriever/postgres/repository.go b/internal/application/repository/retriever/postgres/repository.go
index 773203f03..b974fd6e3 100644
--- a/internal/application/repository/retriever/postgres/repository.go
+++ b/internal/application/repository/retriever/postgres/repository.go
@@ -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))
diff --git a/internal/application/repository/retriever/qdrant/repository.go b/internal/application/repository/retriever/qdrant/repository.go
index 6d660d5ef..7ca72cf91 100644
--- a/internal/application/repository/retriever/qdrant/repository.go
+++ b/internal/application/repository/retriever/qdrant/repository.go
@@ -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
diff --git a/internal/application/service/agent_service.go b/internal/application/service/agent_service.go
index ced67a1aa..4edd15f0d 100644
--- a/internal/application/service/agent_service.go
+++ b/internal/application/service/agent_service.go
@@ -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
+}
diff --git a/internal/application/service/chat_pipline/extract_entity.go b/internal/application/service/chat_pipline/extract_entity.go
index c23f451ee..e1f975712 100644
--- a/internal/application/service/chat_pipline/extract_entity.go
+++ b/internal/application/service/chat_pipline/extract_entity.go
@@ -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,
diff --git a/internal/application/service/chat_pipline/rerank.go b/internal/application/service/chat_pipline/rerank.go
index 7f2769ed7..d23103410 100644
--- a/internal/application/service/chat_pipline/rerank.go
+++ b/internal/application/service/chat_pipline/rerank.go
@@ -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
}
diff --git a/internal/application/service/chat_pipline/search.go b/internal/application/service/chat_pipline/search.go
index 2d249d8a1..75c0402c3 100644
--- a/internal/application/service/chat_pipline/search.go
+++ b/internal/application/service/chat_pipline/search.go
@@ -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 {
diff --git a/internal/application/service/chat_pipline/search_entity.go b/internal/application/service/chat_pipline/search_entity.go
index 94cfa4548..82569e40b 100644
--- a/internal/application/service/chat_pipline/search_entity.go
+++ b/internal/application/service/chat_pipline/search_entity.go
@@ -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()
diff --git a/internal/application/service/chat_pipline/search_parallel.go b/internal/application/service/chat_pipline/search_parallel.go
index 6148016dc..d0f5fa03d 100644
--- a/internal/application/service/chat_pipline/search_parallel.go
+++ b/internal/application/service/chat_pipline/search_parallel.go
@@ -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,
diff --git a/internal/application/service/evaluation.go b/internal/application/service/evaluation.go
index a5b60270f..b71ea63d2 100644
--- a/internal/application/service/evaluation.go
+++ b/internal/application/service/evaluation.go
@@ -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,
)
}
}()
diff --git a/internal/application/service/knowledge.go b/internal/application/service/knowledge.go
index befb20fd2..26aeb4eef 100644
--- a/internal/application/service/knowledge.go
+++ b/internal/application/service/knowledge.go
@@ -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)
+}
+
diff --git a/internal/application/service/knowledgebase.go b/internal/application/service/knowledgebase.go
index d35bdde97..a02be888c 100644
--- a/internal/application/service/knowledgebase.go
+++ b/internal/application/service/knowledgebase.go
@@ -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")
}
diff --git a/internal/application/service/session.go b/internal/application/service/session.go
index 949ad8c13..c277ed57c 100644
--- a/internal/application/service/session.go
+++ b/internal/application/service/session.go
@@ -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
diff --git a/internal/handler/initialization.go b/internal/handler/initialization.go
index 4e35e8560..73b2be758 100644
--- a/internal/handler/initialization.go
+++ b/internal/handler/initialization.go
@@ -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") {
diff --git a/internal/handler/knowledge.go b/internal/handler/knowledge.go
index 460c7f6f9..fa683b775 100644
--- a/internal/handler/knowledge.go
+++ b/internal/handler/knowledge.go
@@ -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,
+ })
+}
diff --git a/internal/handler/session/agent_stream_handler.go b/internal/handler/session/agent_stream_handler.go
index bb39d6ccf..27c1eec10 100644
--- a/internal/handler/session/agent_stream_handler.go
+++ b/internal/handler/session/agent_stream_handler.go
@@ -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
diff --git a/internal/handler/session/qa.go b/internal/handler/session/qa.go
index 9b34c5e38..92dcb56fe 100644
--- a/internal/handler/session/qa.go
+++ b/internal/handler/session/qa.go
@@ -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
diff --git a/internal/handler/session/stream.go b/internal/handler/session/stream.go
index 49d8c9822..9a134cd22 100644
--- a/internal/handler/session/stream.go
+++ b/internal/handler/session/stream.go
@@ -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
}
diff --git a/internal/handler/session/types.go b/internal/handler/session/types.go
index b4aa6f81d..fb0ea153f 100644
--- a/internal/handler/session/types.go
+++ b/internal/handler/session/types.go
@@ -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)
diff --git a/internal/router/router.go b/internal/router/router.go
index 4689366a0..1b9330ce3 100644
--- a/internal/router/router.go
+++ b/internal/router/router.go
@@ -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)
}
}
diff --git a/internal/types/agent.go b/internal/types/agent.go
index fb2464d66..f4529249a 100644
--- a/internal/types/agent.go
+++ b/internal/types/agent.go
@@ -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
diff --git a/internal/types/chat_manage.go b/internal/types/chat_manage.go
index a41a8dcef..576507ba4 100644
--- a/internal/types/chat_manage.go
+++ b/internal/types/chat_manage.go
@@ -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,
diff --git a/internal/types/embedding.go b/internal/types/embedding.go
index ffc320ad6..6fd36b4a0 100644
--- a/internal/types/embedding.go
+++ b/internal/types/embedding.go
@@ -21,6 +21,7 @@ const (
MatchTypeRelationChunk // 关系Chunk匹配类型
MatchTypeGraph
MatchTypeWebSearch // 网络搜索匹配类型
+ MatchTypeDirectLoad // 直接加载匹配类型
)
// IndexInfo contains information about indexed content
diff --git a/internal/types/interfaces/knowledge.go b/internal/types/interfaces/knowledge.go
index 59ced1cc3..4f3edb230 100644
--- a/internal/types/interfaces/knowledge.go
+++ b/internal/types/interfaces/knowledge.go
@@ -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)
}
diff --git a/internal/types/interfaces/knowledgebase.go b/internal/types/interfaces/knowledgebase.go
index e8e91e86a..7ee009f2e 100644
--- a/internal/types/interfaces/knowledgebase.go
+++ b/internal/types/interfaces/knowledgebase.go
@@ -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
diff --git a/internal/types/interfaces/session.go b/internal/types/interfaces/session.go
index e344e6141..b82623ff8 100644
--- a/internal/types/interfaces/session.go
+++ b/internal/types/interfaces/session.go
@@ -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
diff --git a/internal/types/knowledge.go b/internal/types/knowledge.go
index 77627b2e4..ba38d0ed2 100644
--- a/internal/types/knowledge.go
+++ b/internal/types/knowledge.go
@@ -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.
diff --git a/internal/types/retriever.go b/internal/types/retriever.go
index c42e59406..d486cc3d3 100644
--- a/internal/types/retriever.go
+++ b/internal/types/retriever.go
@@ -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
diff --git a/internal/types/search.go b/internal/types/search.go
index f40107c35..95617b3a2 100644
--- a/internal/types/search.go
+++ b/internal/types/search.go
@@ -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