mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-09-24 16:29:01 +08:00
feat: Enhance file upload functionality with progress tracking
- Updated the `uploadKnowledgeFile` function to accept an optional progress callback, enabling real-time upload progress tracking. - Refactored the document upload handling in the menu component to support multiple file uploads with individual progress events. - Removed deprecated document processing state management from the KnowledgeBase view, streamlining the upload process. - Improved user feedback during file uploads by dispatching events for upload start, progress, and completion.
This commit is contained in:
@@ -40,8 +40,8 @@ export function copyKnowledgeBase(data: { source_id: string; target_id?: string
|
||||
}
|
||||
|
||||
// 知识文件 API(基于具体知识库)
|
||||
export function uploadKnowledgeFile(kbId: string, data = {}) {
|
||||
return postUpload(`/api/v1/knowledge-bases/${kbId}/knowledge/file`, data);
|
||||
export function uploadKnowledgeFile(kbId: string, data = {}, onProgress?: (progressEvent: any) => void) {
|
||||
return postUpload(`/api/v1/knowledge-bases/${kbId}/knowledge/file`, data, onProgress);
|
||||
}
|
||||
|
||||
export function createKnowledgeFromURL(kbId: string, data: { url: string; enable_multimodel?: boolean }) {
|
||||
|
||||
@@ -716,38 +716,74 @@ const handleDocFileChange = async (event: Event) => {
|
||||
const totalCount = validFiles.length
|
||||
const failedFiles: Array<{ name: string; reason: string }> = []
|
||||
|
||||
// 显示上传提示
|
||||
if (totalCount > 1) {
|
||||
if (invalidCount > 0) {
|
||||
MessagePlugin.info(t('knowledgeBase.uploadingValidFiles', {
|
||||
valid: totalCount,
|
||||
total: files.length
|
||||
}))
|
||||
} else {
|
||||
MessagePlugin.info(t('knowledgeBase.uploadingMultiple', { total: totalCount }))
|
||||
}
|
||||
}
|
||||
// 为每个文件创建上传任务并发送事件通知
|
||||
const uploadPromises = validFiles.map(async (file) => {
|
||||
const uploadId = `${file.name}_${Date.now()}_${Math.random().toString(36).substr(2, 9)}`
|
||||
let progress = 0
|
||||
let status: 'uploading' | 'success' | 'error' = 'uploading'
|
||||
let error: string | undefined
|
||||
|
||||
// 发送开始上传事件
|
||||
window.dispatchEvent(new CustomEvent('knowledgeFileUploadStart', {
|
||||
detail: {
|
||||
kbId,
|
||||
uploadId,
|
||||
fileName: file.name,
|
||||
file
|
||||
}
|
||||
}))
|
||||
|
||||
for (const file of validFiles) {
|
||||
try {
|
||||
await uploadKnowledgeFile(kbId, { file })
|
||||
await uploadKnowledgeFile(
|
||||
kbId,
|
||||
{ file },
|
||||
(progressEvent: any) => {
|
||||
if (progressEvent.total) {
|
||||
progress = Math.round((progressEvent.loaded * 100) / progressEvent.total)
|
||||
// 发送进度更新事件
|
||||
window.dispatchEvent(new CustomEvent('knowledgeFileUploadProgress', {
|
||||
detail: {
|
||||
kbId,
|
||||
uploadId,
|
||||
progress
|
||||
}
|
||||
}))
|
||||
}
|
||||
}
|
||||
)
|
||||
successCount++
|
||||
status = 'success'
|
||||
progress = 100
|
||||
} catch (error: any) {
|
||||
failCount++
|
||||
let errorReason = error?.error?.message || error?.message || t('knowledgeBase.uploadFailed')
|
||||
if (error?.code === 'duplicate_file' || error?.error?.code === 'duplicate_file') {
|
||||
errorReason = t('knowledgeBase.fileExists')
|
||||
}
|
||||
status = 'error'
|
||||
error = errorReason
|
||||
failedFiles.push({ name: file.name, reason: errorReason })
|
||||
|
||||
// 只在单文件上传时显示详细错误
|
||||
if (totalCount === 1) {
|
||||
MessagePlugin.error(errorReason)
|
||||
} else {
|
||||
// 多文件上传时记录失败信息
|
||||
failedFiles.push({ name: file.name, reason: errorReason })
|
||||
}
|
||||
} finally {
|
||||
// 发送上传完成事件
|
||||
window.dispatchEvent(new CustomEvent('knowledgeFileUploadComplete', {
|
||||
detail: {
|
||||
kbId,
|
||||
uploadId,
|
||||
status,
|
||||
progress,
|
||||
error
|
||||
}
|
||||
}))
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
// 等待所有上传完成
|
||||
await Promise.allSettled(uploadPromises)
|
||||
|
||||
// 显示上传结果
|
||||
if (successCount > 0) {
|
||||
|
||||
@@ -167,12 +167,13 @@ export async function getDown(url: string) {
|
||||
return res
|
||||
}
|
||||
|
||||
export function postUpload(url: string, data = {}) {
|
||||
export function postUpload(url: string, data = {}, onUploadProgress?: (progressEvent: any) => void) {
|
||||
return instance.post(url, data, {
|
||||
headers: {
|
||||
"Content-Type": "multipart/form-data",
|
||||
"X-Request-ID": `${generateRandomString(12)}`,
|
||||
},
|
||||
onUploadProgress,
|
||||
});
|
||||
}
|
||||
|
||||
|
||||
@@ -46,14 +46,6 @@ let knowledgeScroll = ref()
|
||||
let page = 1;
|
||||
let pageSize = 35;
|
||||
|
||||
// 文档处理进度条状态
|
||||
const documentProcessingState = reactive({
|
||||
processingIds: [] as string[],
|
||||
total: 0,
|
||||
completed: 0,
|
||||
failed: 0,
|
||||
pollingInterval: null as ReturnType<typeof setInterval> | null,
|
||||
})
|
||||
const selectedTagId = ref<string>("");
|
||||
const tagList = ref<any[]>([]);
|
||||
const tagLoading = ref(false);
|
||||
@@ -405,18 +397,10 @@ const handleFileUploaded = (event: CustomEvent) => {
|
||||
// 如果上传的文件属于当前知识库,使用 loadKnowledgeFiles 刷新文件列表
|
||||
loadKnowledgeFiles(uploadedKbId);
|
||||
loadTags(uploadedKbId);
|
||||
// 延迟一下,等待文件列表加载完成后再检查处理状态
|
||||
setTimeout(() => {
|
||||
const processingList = cardList.value.filter(item =>
|
||||
item.parse_status === 'pending' || item.parse_status === 'processing'
|
||||
)
|
||||
if (processingList.length > 0) {
|
||||
updateDocumentProcessingState(processingList)
|
||||
}
|
||||
}, 500)
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
// 监听从菜单触发的URL导入事件
|
||||
const handleOpenURLImportDialog = (event: CustomEvent) => {
|
||||
const eventKbId = event.detail.kbId;
|
||||
@@ -434,17 +418,11 @@ onMounted(() => {
|
||||
window.addEventListener('knowledgeFileUploaded', handleFileUploaded as EventListener);
|
||||
// 监听URL导入对话框打开事件
|
||||
window.addEventListener('openURLImportDialog', handleOpenURLImportDialog as EventListener);
|
||||
|
||||
// 恢复文档处理状态
|
||||
if (!isFAQ.value) {
|
||||
restoreDocumentProcessingState()
|
||||
}
|
||||
});
|
||||
|
||||
onUnmounted(() => {
|
||||
window.removeEventListener('knowledgeFileUploaded', handleFileUploaded as EventListener);
|
||||
window.removeEventListener('openURLImportDialog', handleOpenURLImportDialog as EventListener);
|
||||
stopDocumentProcessingPolling()
|
||||
});
|
||||
watch(() => cardList.value, (newValue) => {
|
||||
if (isFAQ.value) return;
|
||||
@@ -460,8 +438,6 @@ watch(() => cardList.value, (newValue) => {
|
||||
updateStatus(analyzeList)
|
||||
}
|
||||
|
||||
// 更新文档处理进度条状态
|
||||
updateDocumentProcessingState(analyzeList)
|
||||
}, { deep: true })
|
||||
type KnowledgeCard = {
|
||||
id: string;
|
||||
@@ -504,167 +480,8 @@ const updateStatus = (analyzeList: KnowledgeCard[]) => {
|
||||
}, 1500);
|
||||
};
|
||||
|
||||
// 更新文档处理进度条状态
|
||||
const updateDocumentProcessingState = (processingList: KnowledgeCard[]) => {
|
||||
const processingIds = processingList.map(item => item.id)
|
||||
const hasChanged = JSON.stringify(processingIds.sort()) !== JSON.stringify(documentProcessingState.processingIds.sort())
|
||||
|
||||
if (hasChanged) {
|
||||
documentProcessingState.processingIds = processingIds
|
||||
documentProcessingState.total = processingIds.length
|
||||
documentProcessingState.completed = 0
|
||||
documentProcessingState.failed = 0
|
||||
|
||||
// 保存到 localStorage
|
||||
saveProcessingIdsToStorage(processingIds)
|
||||
|
||||
// 开始轮询
|
||||
if (processingIds.length > 0) {
|
||||
startDocumentProcessingPolling()
|
||||
} else {
|
||||
stopDocumentProcessingPolling()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 开始轮询文档处理状态
|
||||
const startDocumentProcessingPolling = () => {
|
||||
stopDocumentProcessingPolling()
|
||||
|
||||
if (documentProcessingState.processingIds.length === 0) return
|
||||
|
||||
documentProcessingState.pollingInterval = setInterval(() => {
|
||||
if (documentProcessingState.processingIds.length === 0) {
|
||||
stopDocumentProcessingPolling()
|
||||
return
|
||||
}
|
||||
|
||||
let query = ``
|
||||
documentProcessingState.processingIds.forEach(id => {
|
||||
query += `ids=${id}&`
|
||||
})
|
||||
|
||||
batchQueryKnowledge(query).then((result: any) => {
|
||||
if (result.success && result.data) {
|
||||
const completedIds: string[] = []
|
||||
const failedIds: string[] = []
|
||||
|
||||
;(result.data as KnowledgeCard[]).forEach((item: KnowledgeCard) => {
|
||||
if (item.parse_status === 'completed') {
|
||||
completedIds.push(item.id)
|
||||
documentProcessingState.completed++
|
||||
} else if (item.parse_status === 'failed') {
|
||||
failedIds.push(item.id)
|
||||
documentProcessingState.failed++
|
||||
}
|
||||
})
|
||||
|
||||
// 从处理列表中移除已完成的文档
|
||||
documentProcessingState.processingIds = documentProcessingState.processingIds.filter(
|
||||
id => !completedIds.includes(id) && !failedIds.includes(id)
|
||||
)
|
||||
|
||||
// 如果所有文档都处理完成,停止轮询
|
||||
if (documentProcessingState.processingIds.length === 0) {
|
||||
stopDocumentProcessingPolling()
|
||||
clearProcessingIdsFromStorage()
|
||||
|
||||
// 刷新文件列表
|
||||
if (kbId.value) {
|
||||
loadKnowledgeFiles(kbId.value)
|
||||
}
|
||||
} else {
|
||||
// 更新 localStorage
|
||||
saveProcessingIdsToStorage(documentProcessingState.processingIds)
|
||||
}
|
||||
}
|
||||
}).catch((_err) => {
|
||||
// 错误处理
|
||||
})
|
||||
}, 2000)
|
||||
}
|
||||
|
||||
// 停止轮询文档处理状态
|
||||
const stopDocumentProcessingPolling = () => {
|
||||
if (documentProcessingState.pollingInterval) {
|
||||
clearInterval(documentProcessingState.pollingInterval)
|
||||
documentProcessingState.pollingInterval = null
|
||||
}
|
||||
}
|
||||
|
||||
// localStorage 相关函数
|
||||
const getProcessingIdsStorageKey = () => {
|
||||
return `document_processing_ids_${kbId.value}`
|
||||
}
|
||||
|
||||
const saveProcessingIdsToStorage = (ids: string[]) => {
|
||||
if (!kbId.value) return
|
||||
try {
|
||||
localStorage.setItem(getProcessingIdsStorageKey(), JSON.stringify(ids))
|
||||
} catch (error) {
|
||||
console.error('Failed to save processing IDs to localStorage:', error)
|
||||
}
|
||||
}
|
||||
|
||||
const getProcessingIdsFromStorage = (): string[] => {
|
||||
if (!kbId.value) return []
|
||||
try {
|
||||
const data = localStorage.getItem(getProcessingIdsStorageKey())
|
||||
return data ? JSON.parse(data) : []
|
||||
} catch (error) {
|
||||
console.error('Failed to get processing IDs from localStorage:', error)
|
||||
return []
|
||||
}
|
||||
}
|
||||
|
||||
const clearProcessingIdsFromStorage = () => {
|
||||
if (!kbId.value) return
|
||||
try {
|
||||
localStorage.removeItem(getProcessingIdsStorageKey())
|
||||
} catch (error) {
|
||||
console.error('Failed to clear processing IDs from localStorage:', error)
|
||||
}
|
||||
}
|
||||
|
||||
// 恢复文档处理状态(用于刷新后恢复)
|
||||
const restoreDocumentProcessingState = async () => {
|
||||
if (!kbId.value || isFAQ.value) return
|
||||
|
||||
const savedIds = getProcessingIdsFromStorage()
|
||||
if (savedIds.length === 0) return
|
||||
|
||||
// 检查这些文档是否还在处理中
|
||||
let query = ``
|
||||
savedIds.forEach(id => {
|
||||
query += `ids=${id}&`
|
||||
})
|
||||
|
||||
try {
|
||||
const result: any = await batchQueryKnowledge(query)
|
||||
if (result.success && result.data) {
|
||||
const stillProcessing: string[] = []
|
||||
;(result.data as KnowledgeCard[]).forEach((item: KnowledgeCard) => {
|
||||
if (item.parse_status === 'pending' || item.parse_status === 'processing') {
|
||||
stillProcessing.push(item.id)
|
||||
}
|
||||
})
|
||||
|
||||
if (stillProcessing.length > 0) {
|
||||
documentProcessingState.processingIds = stillProcessing
|
||||
documentProcessingState.total = stillProcessing.length
|
||||
documentProcessingState.completed = 0
|
||||
documentProcessingState.failed = 0
|
||||
saveProcessingIdsToStorage(stillProcessing)
|
||||
startDocumentProcessingPolling()
|
||||
} else {
|
||||
clearProcessingIdsFromStorage()
|
||||
}
|
||||
}
|
||||
} catch (error) {
|
||||
console.error('Failed to restore document processing state:', error)
|
||||
clearProcessingIdsFromStorage()
|
||||
}
|
||||
}
|
||||
|
||||
const closeDoc = () => {
|
||||
isCardDetails.value = false;
|
||||
@@ -709,15 +526,6 @@ const ensureDocumentKbReady = () => {
|
||||
return true;
|
||||
};
|
||||
|
||||
// 关闭文档处理进度条
|
||||
const handleCloseDocumentProgress = () => {
|
||||
stopDocumentProcessingPolling()
|
||||
documentProcessingState.processingIds = []
|
||||
documentProcessingState.total = 0
|
||||
documentProcessingState.completed = 0
|
||||
documentProcessingState.failed = 0
|
||||
clearProcessingIdsFromStorage()
|
||||
}
|
||||
|
||||
const handleDocumentUploadClick = () => {
|
||||
if (!ensureDocumentKbReady()) return;
|
||||
@@ -732,50 +540,92 @@ const resetUploadInput = () => {
|
||||
|
||||
const handleDocumentUpload = async (event: Event) => {
|
||||
const input = event.target as HTMLInputElement;
|
||||
const file = input?.files?.[0];
|
||||
if (!file) return;
|
||||
if (kbFileTypeVerification(file)) {
|
||||
resetUploadInput();
|
||||
return;
|
||||
}
|
||||
const files = input?.files;
|
||||
if (!files || files.length === 0) return;
|
||||
|
||||
if (!kbId.value) {
|
||||
MessagePlugin.error("缺少知识库ID");
|
||||
resetUploadInput();
|
||||
return;
|
||||
}
|
||||
uploading.value = true;
|
||||
try {
|
||||
const responseData: any = await uploadKnowledgeFile(kbId.value, { file });
|
||||
|
||||
// 过滤有效文件
|
||||
const validFiles: File[] = [];
|
||||
for (let i = 0; i < files.length; i++) {
|
||||
const file = files[i];
|
||||
if (!kbFileTypeVerification(file, files.length > 1)) {
|
||||
validFiles.push(file);
|
||||
}
|
||||
}
|
||||
|
||||
if (validFiles.length === 0) {
|
||||
resetUploadInput();
|
||||
return;
|
||||
}
|
||||
|
||||
// 批量上传
|
||||
let successCount = 0;
|
||||
let failCount = 0;
|
||||
const totalCount = validFiles.length;
|
||||
|
||||
for (const file of validFiles) {
|
||||
try {
|
||||
const responseData: any = await uploadKnowledgeFile(kbId.value, { file });
|
||||
const isSuccess = responseData?.success || responseData?.code === 200 || responseData?.status === 'success' || (!responseData?.error && responseData);
|
||||
if (isSuccess) {
|
||||
successCount++;
|
||||
} else {
|
||||
failCount++;
|
||||
let errorMessage = "上传失败!";
|
||||
if (responseData?.error?.message) {
|
||||
errorMessage = responseData.error.message;
|
||||
} else if (responseData?.message) {
|
||||
errorMessage = responseData.message;
|
||||
}
|
||||
if (responseData?.code === 'duplicate_file' || responseData?.error?.code === 'duplicate_file') {
|
||||
errorMessage = "文件已存在";
|
||||
}
|
||||
if (totalCount === 1) {
|
||||
MessagePlugin.error(errorMessage);
|
||||
}
|
||||
}
|
||||
} catch (error: any) {
|
||||
failCount++;
|
||||
let errorMessage = error?.error?.message || error?.message || "上传失败!";
|
||||
if (error?.code === 'duplicate_file') {
|
||||
errorMessage = "文件已存在";
|
||||
}
|
||||
if (totalCount === 1) {
|
||||
MessagePlugin.error(errorMessage);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 显示上传结果
|
||||
if (successCount > 0) {
|
||||
window.dispatchEvent(new CustomEvent('knowledgeFileUploaded', {
|
||||
detail: { kbId: kbId.value }
|
||||
}));
|
||||
const isSuccess = responseData?.success || responseData?.code === 200 || responseData?.status === 'success' || (!responseData?.error && responseData);
|
||||
if (isSuccess) {
|
||||
MessagePlugin.info("上传成功!");
|
||||
} else {
|
||||
let errorMessage = "上传失败!";
|
||||
if (responseData?.error?.message) {
|
||||
errorMessage = responseData.error.message;
|
||||
} else if (responseData?.message) {
|
||||
errorMessage = responseData.message;
|
||||
}
|
||||
if (responseData?.code === 'duplicate_file' || responseData?.error?.code === 'duplicate_file') {
|
||||
errorMessage = "文件已存在";
|
||||
}
|
||||
MessagePlugin.error(errorMessage);
|
||||
}
|
||||
} catch (error: any) {
|
||||
let errorMessage = error?.error?.message || error?.message || "上传失败!";
|
||||
if (error?.code === 'duplicate_file') {
|
||||
errorMessage = "文件已存在";
|
||||
}
|
||||
MessagePlugin.error(errorMessage);
|
||||
} finally {
|
||||
uploading.value = false;
|
||||
resetUploadInput();
|
||||
}
|
||||
|
||||
if (totalCount === 1) {
|
||||
if (successCount === 1) {
|
||||
MessagePlugin.success("上传成功!");
|
||||
}
|
||||
} else {
|
||||
if (failCount === 0) {
|
||||
MessagePlugin.success(`所有文件上传成功(${successCount}个)`);
|
||||
} else if (successCount > 0) {
|
||||
MessagePlugin.warning(`部分文件上传成功(成功:${successCount},失败:${failCount})`);
|
||||
} else {
|
||||
MessagePlugin.error(`所有文件上传失败(${failCount}个)`);
|
||||
}
|
||||
}
|
||||
|
||||
resetUploadInput();
|
||||
};
|
||||
|
||||
|
||||
const handleManualCreate = () => {
|
||||
if (!ensureDocumentKbReady()) return;
|
||||
uiStore.openManualEditor({
|
||||
@@ -1018,45 +868,12 @@ async function createNewSession(value: string): Promise<void> {
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- 文档处理进度条 -->
|
||||
<div v-if="!isFAQ && documentProcessingState.processingIds.length > 0" class="document-processing-progress-bar">
|
||||
<div class="progress-bar-content">
|
||||
<div class="progress-bar-header">
|
||||
<t-icon
|
||||
name="loading"
|
||||
size="16px"
|
||||
class="progress-icon icon-loading"
|
||||
/>
|
||||
<span class="progress-title">
|
||||
{{ $t('knowledgeList.processingDocuments', { count: documentProcessingState.processingIds.length }) }}
|
||||
</span>
|
||||
<span class="progress-count">
|
||||
{{ documentProcessingState.completed + documentProcessingState.failed }}/{{ documentProcessingState.total }}
|
||||
</span>
|
||||
<t-button
|
||||
variant="text"
|
||||
theme="default"
|
||||
size="small"
|
||||
class="progress-close-btn"
|
||||
@click="handleCloseDocumentProgress"
|
||||
>
|
||||
<t-icon name="close" size="14px" />
|
||||
</t-button>
|
||||
</div>
|
||||
<t-progress
|
||||
:percentage="Math.round(((documentProcessingState.completed + documentProcessingState.failed) / documentProcessingState.total) * 100)"
|
||||
:status="documentProcessingState.failed > 0 ? 'error' : 'active'"
|
||||
:label="false"
|
||||
class="progress-bar"
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<input
|
||||
ref="uploadInputRef"
|
||||
type="file"
|
||||
class="document-upload-input"
|
||||
accept=".pdf,.docx,.doc,.txt,.md,.jpg,.jpeg,.png"
|
||||
multiple
|
||||
@change="handleDocumentUpload"
|
||||
/>
|
||||
<div class="knowledge-main">
|
||||
@@ -2106,75 +1923,6 @@ async function createNewSession(value: string): Promise<void> {
|
||||
min-height: 100%;
|
||||
}
|
||||
|
||||
.document-processing-progress-bar {
|
||||
margin-bottom: 16px;
|
||||
background: #fff;
|
||||
border: 1px solid #e7ebf0;
|
||||
border-radius: 8px;
|
||||
padding: 12px 16px;
|
||||
box-shadow: 0 2px 8px rgba(0, 0, 0, 0.04);
|
||||
|
||||
.progress-bar-content {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.progress-bar-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
font-size: 14px;
|
||||
color: #000000e6;
|
||||
|
||||
.progress-icon {
|
||||
flex-shrink: 0;
|
||||
|
||||
&.icon-loading {
|
||||
animation: rotate 1s linear infinite;
|
||||
color: #07c05f;
|
||||
}
|
||||
}
|
||||
|
||||
.progress-title {
|
||||
font-weight: 500;
|
||||
flex: 1;
|
||||
}
|
||||
|
||||
.progress-count {
|
||||
color: #86909c;
|
||||
font-size: 13px;
|
||||
}
|
||||
|
||||
.progress-close-btn {
|
||||
flex-shrink: 0;
|
||||
padding: 4px;
|
||||
margin-left: 8px;
|
||||
}
|
||||
}
|
||||
|
||||
.progress-bar {
|
||||
margin: 0;
|
||||
width: 100%;
|
||||
|
||||
:deep(.t-progress) {
|
||||
width: 100%;
|
||||
}
|
||||
|
||||
:deep(.t-progress__bar) {
|
||||
width: 100%;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@keyframes rotate {
|
||||
from {
|
||||
transform: rotate(0deg);
|
||||
}
|
||||
to {
|
||||
transform: rotate(360deg);
|
||||
}
|
||||
}
|
||||
|
||||
:deep(.del-knowledge) {
|
||||
padding: 0px !important;
|
||||
|
||||
@@ -171,3 +171,22 @@ func (r *chunkRepository) CountChunksByKnowledgeBaseID(ctx context.Context, tena
|
||||
Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (r *chunkRepository) DeleteChunksByChunkIndexRange(ctx context.Context, tenantID uint, knowledgeID string, startChunkIndex int, endChunkIndex int) ([]*types.Chunk, error) {
|
||||
var chunks []*types.Chunk
|
||||
// 先查询要删除的chunks
|
||||
if err := r.db.WithContext(ctx).
|
||||
Where("tenant_id = ? AND knowledge_id = ? AND chunk_index >= ? AND chunk_index <= ? AND deleted_at IS NULL", tenantID, knowledgeID, startChunkIndex, endChunkIndex).
|
||||
Find(&chunks).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 然后删除它们
|
||||
if len(chunks) > 0 {
|
||||
if err := r.db.WithContext(ctx).
|
||||
Where("tenant_id = ? AND knowledge_id = ? AND chunk_index >= ? AND chunk_index <= ? AND deleted_at IS NULL", tenantID, knowledgeID, startChunkIndex, endChunkIndex).
|
||||
Delete(&types.Chunk{}).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return chunks, nil
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"regexp"
|
||||
"runtime"
|
||||
"slices"
|
||||
"sort"
|
||||
"strings"
|
||||
@@ -276,10 +277,40 @@ func (s *knowledgeService) CreateKnowledgeFromFile(ctx context.Context,
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Process document asynchronously
|
||||
logger.Info(ctx, "Starting asynchronous document processing")
|
||||
newCtx := logger.CloneContext(ctx)
|
||||
go s.processDocument(newCtx, kb, knowledge, file, kb.VLMConfig.Enabled)
|
||||
// Enqueue document processing task to Asynq
|
||||
logger.Info(ctx, "Enqueuing document processing task to Asynq")
|
||||
enableMultimodelValue := false
|
||||
if enableMultimodel != nil {
|
||||
enableMultimodelValue = *enableMultimodel
|
||||
} else {
|
||||
enableMultimodelValue = kb.VLMConfig.Enabled
|
||||
}
|
||||
|
||||
taskPayload := types.DocumentProcessPayload{
|
||||
TenantID: tenantID,
|
||||
KnowledgeID: knowledge.ID,
|
||||
KnowledgeBaseID: kbID,
|
||||
FilePath: filePath,
|
||||
FileName: safeFilename,
|
||||
FileType: getFileType(safeFilename),
|
||||
EnableMultimodel: enableMultimodelValue,
|
||||
}
|
||||
|
||||
payloadBytes, err := json.Marshal(taskPayload)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to marshal document process task payload: %v", err)
|
||||
// 即使入队失败,也返回knowledge,因为文件已保存
|
||||
return knowledge, nil
|
||||
}
|
||||
|
||||
task := asynq.NewTask(types.TypeDocumentProcess, payloadBytes, asynq.Queue("default"))
|
||||
info, err := s.task.Enqueue(task)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to enqueue document process task: %v", err)
|
||||
// 即使入队失败,也返回knowledge,因为文件已保存
|
||||
return knowledge, nil
|
||||
}
|
||||
logger.Infof(ctx, "Enqueued document process task: id=%s queue=%s knowledge_id=%s", info.ID, info.Queue, knowledge.ID)
|
||||
|
||||
logger.Infof(ctx, "Knowledge from file created successfully, ID: %s", knowledge.ID)
|
||||
return knowledge, nil
|
||||
@@ -362,10 +393,36 @@ func (s *knowledgeService) CreateKnowledgeFromURL(ctx context.Context,
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Process URL asynchronously
|
||||
logger.Info(ctx, "Starting asynchronous URL processing")
|
||||
newCtx := logger.CloneContext(ctx)
|
||||
go s.processDocumentFromURL(newCtx, kb, knowledge, url, kb.VLMConfig.Enabled)
|
||||
// Enqueue URL processing task to Asynq
|
||||
logger.Info(ctx, "Enqueuing URL processing task to Asynq")
|
||||
enableMultimodelValue := false
|
||||
if enableMultimodel != nil {
|
||||
enableMultimodelValue = *enableMultimodel
|
||||
} else {
|
||||
enableMultimodelValue = kb.VLMConfig.Enabled
|
||||
}
|
||||
|
||||
taskPayload := types.DocumentProcessPayload{
|
||||
TenantID: tenantID,
|
||||
KnowledgeID: knowledge.ID,
|
||||
KnowledgeBaseID: kbID,
|
||||
URL: url,
|
||||
EnableMultimodel: enableMultimodelValue,
|
||||
}
|
||||
|
||||
payloadBytes, err := json.Marshal(taskPayload)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to marshal URL process task payload: %v", err)
|
||||
return knowledge, nil
|
||||
}
|
||||
|
||||
task := asynq.NewTask(types.TypeDocumentProcess, payloadBytes, asynq.Queue("default"))
|
||||
info, err := s.task.Enqueue(task)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to enqueue URL process task: %v", err)
|
||||
return knowledge, nil
|
||||
}
|
||||
logger.Infof(ctx, "Enqueued URL process task: id=%s queue=%s knowledge_id=%s", info.ID, info.Queue, knowledge.ID)
|
||||
|
||||
logger.Infof(ctx, "Knowledge from URL created successfully, ID: %s", knowledge.ID)
|
||||
return knowledge, nil
|
||||
@@ -532,8 +589,31 @@ func (s *knowledgeService) createKnowledgeFromPassageInternal(ctx context.Contex
|
||||
s.processDocumentFromPassage(ctx, kb, knowledge, safePassages)
|
||||
logger.Infof(ctx, "Knowledge from passage created successfully (sync), ID: %s", knowledge.ID)
|
||||
} else {
|
||||
logger.Info(ctx, "Starting asynchronous passage processing")
|
||||
go s.processDocumentFromPassage(ctx, kb, knowledge, safePassages)
|
||||
// Enqueue passage processing task to Asynq
|
||||
logger.Info(ctx, "Enqueuing passage processing task to Asynq")
|
||||
tenantID := ctx.Value(types.TenantIDContextKey).(uint)
|
||||
taskPayload := types.DocumentProcessPayload{
|
||||
TenantID: tenantID,
|
||||
KnowledgeID: knowledge.ID,
|
||||
KnowledgeBaseID: kbID,
|
||||
Passages: safePassages,
|
||||
EnableMultimodel: false, // 文本段落不支持多模态
|
||||
}
|
||||
|
||||
payloadBytes, err := json.Marshal(taskPayload)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to marshal passage process task payload: %v", err)
|
||||
// 即使入队失败,也返回knowledge
|
||||
return knowledge, nil
|
||||
}
|
||||
|
||||
task := asynq.NewTask(types.TypeDocumentProcess, payloadBytes, asynq.Queue("default"))
|
||||
info, err := s.task.Enqueue(task)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to enqueue passage process task: %v", err)
|
||||
return knowledge, nil
|
||||
}
|
||||
logger.Infof(ctx, "Enqueued passage process task: id=%s queue=%s knowledge_id=%s", info.ID, info.Queue, knowledge.ID)
|
||||
logger.Infof(ctx, "Knowledge from passage created successfully, ID: %s", knowledge.ID)
|
||||
}
|
||||
return knowledge, nil
|
||||
@@ -1046,6 +1126,36 @@ func (s *knowledgeService) processChunks(ctx context.Context,
|
||||
return
|
||||
}
|
||||
|
||||
// 幂等性处理:清理旧的chunks和索引数据,避免重复数据
|
||||
logger.Infof(ctx, "Cleaning up existing chunks and index data for knowledge: %s", knowledge.ID)
|
||||
|
||||
// 删除旧的chunks
|
||||
if err := s.chunkService.DeleteChunksByKnowledgeID(ctx, knowledge.ID); err != nil {
|
||||
logger.Warnf(ctx, "Failed to delete existing chunks (may not exist): %v", err)
|
||||
// 不返回错误,继续处理(可能没有旧数据)
|
||||
}
|
||||
|
||||
// 删除旧的索引数据
|
||||
tenantInfo := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
retrieveEngine, err := retriever.NewCompositeRetrieveEngine(s.retrieveEngine, tenantInfo.RetrieverEngines.Engines)
|
||||
if err == nil {
|
||||
if err := retrieveEngine.DeleteByKnowledgeIDList(ctx, []string{knowledge.ID}, embeddingModel.GetDimensions()); err != nil {
|
||||
logger.Warnf(ctx, "Failed to delete existing index data (may not exist): %v", err)
|
||||
// 不返回错误,继续处理(可能没有旧数据)
|
||||
} else {
|
||||
logger.Infof(ctx, "Successfully deleted existing index data for knowledge: %s", knowledge.ID)
|
||||
}
|
||||
}
|
||||
|
||||
// 删除知识图谱数据(如果存在)
|
||||
namespace := types.NameSpace{KnowledgeBase: knowledge.KnowledgeBaseID, Knowledge: knowledge.ID}
|
||||
if err := s.graphEngine.DelGraph(ctx, []types.NameSpace{namespace}); err != nil {
|
||||
logger.Warnf(ctx, "Failed to delete existing graph data (may not exist): %v", err)
|
||||
// 不返回错误,继续处理
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Cleanup completed, starting to process new chunks")
|
||||
|
||||
// Generate document summary - 只使用文本类型的 Chunk
|
||||
chatModel, err := s.modelService.GetChatModel(ctx, kb.SummaryModelID)
|
||||
if err != nil {
|
||||
@@ -1243,16 +1353,6 @@ func (s *knowledgeService) processChunks(ctx context.Context,
|
||||
})
|
||||
|
||||
// Initialize retrieval engine
|
||||
tenantInfo := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
retrieveEngine, err := retriever.NewCompositeRetrieveEngine(s.retrieveEngine, tenantInfo.RetrieverEngines.Engines)
|
||||
if err != nil {
|
||||
knowledge.ParseStatus = "failed"
|
||||
knowledge.ErrorMessage = err.Error()
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
s.repo.UpdateKnowledge(ctx, knowledge)
|
||||
span.RecordError(err)
|
||||
return
|
||||
}
|
||||
|
||||
// Calculate storage size required for embeddings
|
||||
span.AddEvent("estimate storage size")
|
||||
@@ -2083,7 +2183,8 @@ func (s *knowledgeService) UpsertFAQEntries(ctx context.Context,
|
||||
}
|
||||
|
||||
// 验证知识库是否存在且有效
|
||||
if _, err := s.validateFAQKnowledgeBase(ctx, kbID); err != nil {
|
||||
kb, err := s.validateFAQKnowledgeBase(ctx, kbID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
@@ -2099,11 +2200,6 @@ func (s *knowledgeService) UpsertFAQEntries(ctx context.Context,
|
||||
return "", werrors.NewBadRequestError(fmt.Sprintf("该知识库已有导入任务正在进行中(任务ID: %s),请等待完成后再试", runningKnowledge.ID))
|
||||
}
|
||||
|
||||
// 确保FAQ Knowledge存在
|
||||
kb, err := s.validateFAQKnowledgeBase(ctx, kbID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
faqKnowledge, err := s.ensureFAQKnowledge(ctx, tenantID, kb)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to ensure FAQ knowledge: %w", err)
|
||||
@@ -2111,63 +2207,69 @@ func (s *knowledgeService) UpsertFAQEntries(ctx context.Context,
|
||||
|
||||
// 初始化导入任务状态到Knowledge表
|
||||
taskID := faqKnowledge.ID // 使用Knowledge ID作为taskID
|
||||
if err := s.updateFAQImportStatus(ctx, taskID, types.FAQImportStatusPending, 0, len(payload.Entries), 0, ""); err != nil {
|
||||
|
||||
// 从Knowledge Metadata中读取当前的NextChunkIndex(自增计数器)
|
||||
// 如果不存在,则初始化为1
|
||||
existingMeta, _ := types.ParseFAQImportMetadata(faqKnowledge)
|
||||
nextChunkIndex := 1
|
||||
if existingMeta != nil && existingMeta.NextChunkIndex > 0 {
|
||||
nextChunkIndex = existingMeta.NextChunkIndex
|
||||
}
|
||||
startChunkIndex := nextChunkIndex
|
||||
endChunkIndex := startChunkIndex + len(payload.Entries) - 1
|
||||
|
||||
// 更新NextChunkIndex为下一个要分配的值
|
||||
newNextChunkIndex := nextChunkIndex + len(payload.Entries)
|
||||
|
||||
// 更新Metadata,包含区间信息和新的NextChunkIndex
|
||||
importMeta := &types.FAQImportMetadata{
|
||||
ImportProgress: 0,
|
||||
ImportTotal: len(payload.Entries),
|
||||
ImportProcessed: 0,
|
||||
NextChunkIndex: newNextChunkIndex,
|
||||
}
|
||||
|
||||
if err := s.updateFAQImportStatusWithRanges(ctx, taskID, types.FAQImportStatusPending, 0, len(payload.Entries), 0, "", importMeta.NextChunkIndex); err != nil {
|
||||
logger.Errorf(ctx, "Failed to initialize FAQ import task status: %v", err)
|
||||
return "", fmt.Errorf("failed to initialize task: %w", err)
|
||||
}
|
||||
bgCtx := logger.CloneContext(ctx)
|
||||
// 在后台goroutine中执行导入
|
||||
go func() {
|
||||
// 使用独立的context,不受HTTP请求context影响
|
||||
// 设置较长的超时时间(2小时)
|
||||
bgCtx, cancel := context.WithTimeout(bgCtx, 2*time.Hour)
|
||||
defer cancel()
|
||||
|
||||
logger.Infof(bgCtx, "Starting FAQ import task: %s, total entries: %d", taskID, len(payload.Entries))
|
||||
logger.Infof(ctx, "Allocated ChunkIndex range [%d, %d] for FAQ import task %s, next ChunkIndex will be %d",
|
||||
startChunkIndex, endChunkIndex, taskID, newNextChunkIndex)
|
||||
|
||||
// 更新任务状态为运行中
|
||||
if err := s.updateFAQImportStatus(bgCtx, taskID, types.FAQImportStatusRunning, 0, len(payload.Entries), 0, ""); err != nil {
|
||||
logger.Errorf(bgCtx, "Failed to update task status to running: %v", err)
|
||||
}
|
||||
// Enqueue FAQ import task to Asynq
|
||||
logger.Info(ctx, "Enqueuing FAQ import task to Asynq")
|
||||
taskPayload := types.FAQImportPayload{
|
||||
TenantID: tenantID,
|
||||
TaskID: taskID,
|
||||
KBID: kbID,
|
||||
KnowledgeID: faqKnowledge.ID,
|
||||
Entries: payload.Entries,
|
||||
Mode: payload.Mode,
|
||||
StartChunkIndex: startChunkIndex,
|
||||
EndChunkIndex: endChunkIndex,
|
||||
}
|
||||
|
||||
// 执行实际导入
|
||||
if err := s.executeFAQImport(bgCtx, taskID, kbID, payload, tenantID); err != nil {
|
||||
logger.Errorf(bgCtx, "FAQ import task failed: %s, error: %v", taskID, err)
|
||||
// 获取当前进度
|
||||
tenantID := bgCtx.Value(types.TenantIDContextKey).(uint)
|
||||
knowledge, _ := s.repo.GetKnowledgeByID(bgCtx, tenantID, taskID)
|
||||
total := len(payload.Entries)
|
||||
processed := 0
|
||||
if knowledge != nil {
|
||||
importMeta, _ := types.ParseFAQImportMetadata(knowledge)
|
||||
if importMeta != nil {
|
||||
total = importMeta.ImportTotal
|
||||
processed = importMeta.ImportProcessed
|
||||
}
|
||||
}
|
||||
if updateErr := s.updateFAQImportStatus(bgCtx, taskID, types.FAQImportStatusFailed, 0, total, processed, err.Error()); updateErr != nil {
|
||||
logger.Errorf(bgCtx, "Failed to update task status to failed: %v", updateErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
payloadBytes, err := json.Marshal(taskPayload)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to marshal FAQ import task payload: %v", err)
|
||||
return "", fmt.Errorf("failed to marshal task payload: %w", err)
|
||||
}
|
||||
|
||||
// 任务成功完成
|
||||
logger.Infof(bgCtx, "FAQ import task completed: %s", taskID)
|
||||
if err := s.updateFAQImportStatus(bgCtx, taskID, types.FAQImportStatusSuccess, 100, len(payload.Entries), len(payload.Entries), ""); err != nil {
|
||||
logger.Errorf(bgCtx, "Failed to update task status to success: %v", err)
|
||||
}
|
||||
}()
|
||||
task := asynq.NewTask(types.TypeFAQImport, payloadBytes, asynq.Queue("default"), asynq.MaxRetry(5))
|
||||
info, err := s.task.Enqueue(task)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to enqueue FAQ import task: %v", err)
|
||||
return "", fmt.Errorf("failed to enqueue task: %w", err)
|
||||
}
|
||||
logger.Infof(ctx, "Enqueued FAQ import task: id=%s queue=%s task_id=%s", info.ID, info.Queue, taskID)
|
||||
|
||||
return taskID, nil
|
||||
}
|
||||
|
||||
// executeFAQImport 执行实际的FAQ导入逻辑
|
||||
func (s *knowledgeService) executeFAQImport(ctx context.Context, taskID string, kbID string,
|
||||
payload *types.FAQBatchUpsertPayload, tenantID uint) (err error) {
|
||||
// 用于记录所有已创建的chunks,用于失败时回滚
|
||||
var createdChunks []*types.Chunk
|
||||
// 用于记录已索引的chunks,需要清理索引数据
|
||||
var indexedChunks []*types.Chunk
|
||||
payload *types.FAQBatchUpsertPayload, tenantID uint, startChunkIndex int, processed int, totalEntries int) (err error) {
|
||||
// 保存知识库和embedding模型信息,用于清理索引
|
||||
var kb *types.KnowledgeBase
|
||||
var embeddingModel embedding.Embedder
|
||||
@@ -2176,47 +2278,12 @@ func (s *knowledgeService) executeFAQImport(ctx context.Context, taskID string,
|
||||
defer func() {
|
||||
// 捕获panic
|
||||
if r := recover(); r != nil {
|
||||
logger.Errorf(ctx, "FAQ import task %s panicked: %v", taskID, r)
|
||||
buf := make([]byte, 8192)
|
||||
n := runtime.Stack(buf, false)
|
||||
stack := string(buf[:n])
|
||||
logger.Errorf(ctx, "FAQ import task %s panicked: %v\n%s", taskID, r, stack)
|
||||
err = fmt.Errorf("panic during FAQ import: %v", r)
|
||||
}
|
||||
if err != nil && len(createdChunks) > 0 {
|
||||
logger.Warnf(ctx, "FAQ import task %s failed (error: %v), rolling back %d created chunks and their indices", taskID, err, len(createdChunks))
|
||||
|
||||
// 清理索引数据(如果有已索引的chunks)
|
||||
if len(indexedChunks) > 0 && kb != nil && embeddingModel != nil {
|
||||
chunkIDs := make([]string, 0, len(indexedChunks))
|
||||
for _, chunk := range indexedChunks {
|
||||
chunkIDs = append(chunkIDs, chunk.ID)
|
||||
}
|
||||
// 从context获取tenant信息
|
||||
tenantInfo := ctx.Value(types.TenantInfoContextKey)
|
||||
if tenantInfo != nil {
|
||||
if tenant, ok := tenantInfo.(*types.Tenant); ok {
|
||||
retrieveEngine, engineErr := retriever.NewCompositeRetrieveEngine(s.retrieveEngine, tenant.RetrieverEngines.Engines)
|
||||
if engineErr == nil {
|
||||
if delIndexErr := retrieveEngine.DeleteByChunkIDList(ctx, chunkIDs, embeddingModel.GetDimensions()); delIndexErr != nil {
|
||||
logger.Errorf(ctx, "Failed to delete indices for %d chunks during rollback: %v", len(chunkIDs), delIndexErr)
|
||||
} else {
|
||||
logger.Debugf(ctx, "Successfully deleted indices for %d chunks", len(chunkIDs))
|
||||
}
|
||||
} else {
|
||||
logger.Errorf(ctx, "Failed to create retrieve engine during rollback: %v", engineErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 批量删除chunks,提高回滚效率
|
||||
chunkIDs := make([]string, 0, len(createdChunks))
|
||||
for _, chunk := range createdChunks {
|
||||
chunkIDs = append(chunkIDs, chunk.ID)
|
||||
}
|
||||
if delErr := s.chunkService.DeleteChunks(ctx, chunkIDs); delErr != nil {
|
||||
logger.Errorf(ctx, "Failed to delete %d chunks during rollback: %v", len(chunkIDs), delErr)
|
||||
} else {
|
||||
logger.Debugf(ctx, "Successfully rolled back %d chunks", len(chunkIDs))
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
kb, err = s.validateFAQKnowledgeBase(ctx, kbID)
|
||||
@@ -2243,21 +2310,6 @@ func (s *knowledgeService) executeFAQImport(ctx context.Context, taskID string,
|
||||
}
|
||||
}
|
||||
|
||||
startIndex := 0
|
||||
if payload.Mode == types.FAQBatchModeAppend {
|
||||
_, total, err := s.chunkRepo.ListPagedChunksByKnowledgeID(ctx,
|
||||
tenantID,
|
||||
faqKnowledge.ID,
|
||||
&types.Pagination{Page: 1, PageSize: 1},
|
||||
[]types.ChunkType{types.ChunkTypeFAQ},
|
||||
"",
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
startIndex = int(total)
|
||||
}
|
||||
|
||||
// 获取索引模式
|
||||
indexMode := types.FAQIndexModeQuestionOnly
|
||||
if kb.FAQConfig != nil && kb.FAQConfig.IndexMode != "" {
|
||||
@@ -2265,17 +2317,18 @@ func (s *knowledgeService) executeFAQImport(ctx context.Context, taskID string,
|
||||
}
|
||||
|
||||
// 分批处理
|
||||
totalEntries := len(payload.Entries)
|
||||
processed := 0
|
||||
batchStartTime := time.Now()
|
||||
remainingEntries := len(payload.Entries)
|
||||
totalStartTime := time.Now()
|
||||
// 保存初始的processed值,用于计算ChunkIndex(i+idx是剩余条目中的索引,加上初始processed得到原始位置)
|
||||
initialProcessed := processed
|
||||
|
||||
logger.Infof(ctx, "FAQ import task %s: starting batch processing, total entries: %d, batch size: %d", taskID, totalEntries, faqImportBatchSize)
|
||||
logger.Infof(ctx, "FAQ import task %s: starting batch processing, remaining entries: %d, total entries: %d, batch size: %d", taskID, remainingEntries, totalEntries, faqImportBatchSize)
|
||||
|
||||
for i := 0; i < totalEntries; i += faqImportBatchSize {
|
||||
batchStartTime = time.Now()
|
||||
for i := 0; i < remainingEntries; i += faqImportBatchSize {
|
||||
batchStartTime := time.Now()
|
||||
end := i + faqImportBatchSize
|
||||
if end > totalEntries {
|
||||
end = totalEntries
|
||||
if end > remainingEntries {
|
||||
end = remainingEntries
|
||||
}
|
||||
|
||||
batch := payload.Entries[i:end]
|
||||
@@ -2284,6 +2337,7 @@ func (s *knowledgeService) executeFAQImport(ctx context.Context, taskID string,
|
||||
// 构建chunks
|
||||
buildStartTime := time.Now()
|
||||
chunks := make([]*types.Chunk, 0, len(batch))
|
||||
chunkIds := make([]string, 0, len(batch))
|
||||
for idx, entry := range batch {
|
||||
meta, err := sanitizeFAQEntryPayload(&entry)
|
||||
if err != nil {
|
||||
@@ -2297,13 +2351,15 @@ func (s *knowledgeService) executeFAQImport(ctx context.Context, taskID string,
|
||||
if entry.IsEnabled != nil {
|
||||
isEnabled = *entry.IsEnabled
|
||||
}
|
||||
// ChunkIndex计算:startChunkIndex + (i+idx) + initialProcessed
|
||||
// i+idx是剩余条目中的索引,加上initialProcessed得到原始条目中的位置
|
||||
chunk := &types.Chunk{
|
||||
ID: uuid.New().String(),
|
||||
TenantID: tenantID,
|
||||
KnowledgeID: faqKnowledge.ID,
|
||||
KnowledgeBaseID: kb.ID,
|
||||
Content: buildFAQChunkContent(meta, indexMode),
|
||||
ChunkIndex: startIndex + i + idx + 1,
|
||||
ChunkIndex: startChunkIndex + i + idx + initialProcessed,
|
||||
IsEnabled: isEnabled,
|
||||
ChunkType: types.ChunkTypeFAQ,
|
||||
TagID: entry.TagID,
|
||||
@@ -2312,10 +2368,11 @@ func (s *knowledgeService) executeFAQImport(ctx context.Context, taskID string,
|
||||
return fmt.Errorf("failed to set FAQ metadata: %w", err)
|
||||
}
|
||||
chunks = append(chunks, chunk)
|
||||
chunkIds = append(chunkIds, chunk.ID)
|
||||
}
|
||||
buildDuration := time.Since(buildStartTime)
|
||||
logger.Debugf(ctx, "FAQ import task %s: batch %d-%d built %d chunks in %v", taskID, i+1, end, len(chunks), buildDuration)
|
||||
|
||||
logger.Debugf(ctx, "FAQ import task %s: batch %d-%d built %d chunks in %v, chunk IDs: %v",
|
||||
taskID, i+1, end, len(chunks), buildDuration, chunkIds)
|
||||
// 创建chunks
|
||||
createStartTime := time.Now()
|
||||
if err := s.chunkService.CreateChunks(ctx, chunks); err != nil {
|
||||
@@ -2323,8 +2380,6 @@ func (s *knowledgeService) executeFAQImport(ctx context.Context, taskID string,
|
||||
}
|
||||
createDuration := time.Since(createStartTime)
|
||||
logger.Infof(ctx, "FAQ import task %s: batch %d-%d created %d chunks in %v", taskID, i+1, end, len(chunks), createDuration)
|
||||
// 记录已创建的chunks(用于失败时回滚)
|
||||
createdChunks = append(createdChunks, chunks...)
|
||||
|
||||
// 索引chunks
|
||||
indexStartTime := time.Now()
|
||||
@@ -2334,30 +2389,23 @@ func (s *knowledgeService) executeFAQImport(ctx context.Context, taskID string,
|
||||
}
|
||||
indexDuration := time.Since(indexStartTime)
|
||||
logger.Infof(ctx, "FAQ import task %s: batch %d-%d indexed %d chunks in %v", taskID, i+1, end, len(chunks), indexDuration)
|
||||
// 记录已成功索引的chunks(用于失败时清理索引数据)
|
||||
indexedChunks = append(indexedChunks, chunks...)
|
||||
|
||||
processed += len(batch)
|
||||
processed = processed + len(batch)
|
||||
|
||||
// 更新任务进度
|
||||
progressUpdateStartTime := time.Now()
|
||||
progress := int(float64(processed) / float64(totalEntries) * 100)
|
||||
if err := s.updateFAQImportStatus(ctx, taskID, types.FAQImportStatusRunning, progress, totalEntries, processed, ""); err != nil {
|
||||
if err := s.updateFAQImportStatus(ctx, taskID, types.FAQImportStatusProcessing, progress, totalEntries, processed, ""); err != nil {
|
||||
logger.Errorf(ctx, "Failed to update task progress: %v", err)
|
||||
}
|
||||
progressUpdateDuration := time.Since(progressUpdateStartTime)
|
||||
if progressUpdateDuration > 100*time.Millisecond {
|
||||
logger.Warnf(ctx, "FAQ import task %s: progress update took %v (may be slow)", taskID, progressUpdateDuration)
|
||||
}
|
||||
|
||||
batchDuration := time.Since(batchStartTime)
|
||||
logger.Infof(ctx, "FAQ import task %s: batch %d-%d completed in %v (build: %v, create: %v, index: %v, progress: %v), total progress: %d/%d (%d%%)",
|
||||
taskID, i+1, end, batchDuration, buildDuration, createDuration, indexDuration, progressUpdateDuration, processed, totalEntries, progress)
|
||||
logger.Infof(ctx, "FAQ import task %s: batch %d-%d completed in %v (build: %v, create: %v, index: %v), total progress: %d/%d (%d%%)",
|
||||
taskID, i+1, end, batchDuration, buildDuration, createDuration, indexDuration, processed, totalEntries, progress)
|
||||
}
|
||||
|
||||
totalDuration := time.Since(batchStartTime)
|
||||
logger.Infof(ctx, "FAQ import task %s: all batches completed, total: %d entries in %v, avg: %v per entry",
|
||||
taskID, totalEntries, totalDuration, totalDuration/time.Duration(totalEntries))
|
||||
totalDuration := time.Since(totalStartTime)
|
||||
logger.Infof(ctx, "FAQ import task %s: all batches completed, processed: %d entries in %v, avg: %v per entry",
|
||||
taskID, remainingEntries, totalDuration, totalDuration/time.Duration(remainingEntries))
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -2938,8 +2986,14 @@ func (s *knowledgeService) ensureFAQKnowledge(ctx context.Context, tenantID uint
|
||||
return knowledge, nil
|
||||
}
|
||||
|
||||
// updateFAQImportStatus 更新FAQ Knowledge的导入任务状态
|
||||
func (s *knowledgeService) updateFAQImportStatus(ctx context.Context, knowledgeID string, status types.FAQImportTaskStatus, progress, total, processed int, errorMsg string) error {
|
||||
func (s *knowledgeService) updateFAQImportStatus(ctx context.Context, knowledgeID string, status types.FAQImportTaskStatus,
|
||||
progress, total, processed int, errorMsg string) error {
|
||||
return s.updateFAQImportStatusWithRanges(ctx, knowledgeID, status, progress, total, processed, errorMsg, 0)
|
||||
}
|
||||
|
||||
// updateFAQImportStatusWithRanges 更新FAQ Knowledge的导入任务状态,包含NextChunkIndex
|
||||
func (s *knowledgeService) updateFAQImportStatusWithRanges(ctx context.Context, knowledgeID string, status types.FAQImportTaskStatus,
|
||||
progress, total, processed int, errorMsg string, nextChunkIndex int) error {
|
||||
tenantID := ctx.Value(types.TenantIDContextKey).(uint)
|
||||
knowledge, err := s.repo.GetKnowledgeByID(ctx, tenantID, knowledgeID)
|
||||
if err != nil {
|
||||
@@ -2947,30 +3001,29 @@ func (s *knowledgeService) updateFAQImportStatus(ctx context.Context, knowledgeI
|
||||
}
|
||||
|
||||
// 更新ParseStatus:将FAQImportTaskStatus映射到ParseStatus
|
||||
// pending -> "pending", running -> "processing", success -> "completed", failed -> "failed"
|
||||
parseStatus := string(status)
|
||||
if status == types.FAQImportStatusRunning {
|
||||
parseStatus = "processing" // 使用"processing"以兼容文档类型的ParseStatus
|
||||
} else if status == types.FAQImportStatusSuccess {
|
||||
parseStatus = "completed"
|
||||
}
|
||||
knowledge.ParseStatus = parseStatus
|
||||
|
||||
knowledge.ParseStatus = string(status)
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
|
||||
// 更新ErrorMessage
|
||||
if errorMsg != "" {
|
||||
knowledge.ErrorMessage = errorMsg
|
||||
} else if status == types.FAQImportStatusSuccess {
|
||||
knowledge.ErrorMessage = "" // 成功时清空错误信息
|
||||
meta, err := types.ParseFAQImportMetadata(knowledge)
|
||||
if err != nil || meta == nil {
|
||||
meta = &types.FAQImportMetadata{}
|
||||
}
|
||||
|
||||
// 更新Metadata中的导入进度信息
|
||||
importMeta := &types.FAQImportMetadata{
|
||||
ImportProgress: progress,
|
||||
ImportTotal: total,
|
||||
ImportProcessed: processed,
|
||||
// 更新ErrorMessage
|
||||
knowledge.ErrorMessage = errorMsg
|
||||
if status == types.FAQImportStatusCompleted {
|
||||
knowledge.ErrorMessage = ""
|
||||
}
|
||||
metaJSON, err := importMeta.ToJSON()
|
||||
|
||||
// 更新Metadata中的导入进度信息,保留已有的ChunkIndexRanges和NextChunkIndex
|
||||
meta.ImportProgress = progress
|
||||
meta.ImportTotal = total
|
||||
meta.ImportProcessed = processed
|
||||
if nextChunkIndex > 0 {
|
||||
meta.NextChunkIndex = nextChunkIndex
|
||||
}
|
||||
metaJSON, err := meta.ToJSON()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal import metadata: %w", err)
|
||||
}
|
||||
@@ -3224,7 +3277,8 @@ func (s *knowledgeService) indexFAQChunks(ctx context.Context,
|
||||
return err
|
||||
}
|
||||
batchIndexDuration := time.Since(batchIndexStartTime)
|
||||
logger.Debugf(ctx, "indexFAQChunks: batch indexed %d index info entries in %v (avg: %v per entry)", len(indexInfo), batchIndexDuration, batchIndexDuration/time.Duration(len(indexInfo)))
|
||||
logger.Debugf(ctx, "indexFAQChunks: batch indexed %d index info entries in %v (avg: %v per entry)",
|
||||
len(indexInfo), batchIndexDuration, batchIndexDuration/time.Duration(len(indexInfo)))
|
||||
|
||||
if adjustStorage && size > 0 {
|
||||
adjustStartTime := time.Now()
|
||||
@@ -3513,3 +3567,359 @@ func IsImageType(fileType string) bool {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// ProcessDocument handles Asynq document processing tasks
|
||||
func (s *knowledgeService) ProcessDocument(ctx context.Context, t *asynq.Task) error {
|
||||
var payload types.DocumentProcessPayload
|
||||
if err := json.Unmarshal(t.Payload(), &payload); err != nil {
|
||||
logger.Errorf(ctx, "failed to unmarshal document process task payload: %v", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
ctx = logger.WithRequestID(ctx, payload.RequestId)
|
||||
ctx = logger.WithField(ctx, "document_process", payload.KnowledgeID)
|
||||
ctx = context.WithValue(ctx, types.TenantIDContextKey, payload.TenantID)
|
||||
tenantInfo, err := s.tenantRepo.GetTenantByID(ctx, payload.TenantID)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "failed to get tenant: %v", err)
|
||||
return nil
|
||||
}
|
||||
ctx = context.WithValue(ctx, types.TenantInfoContextKey, tenantInfo)
|
||||
|
||||
logger.Infof(ctx, "Processing document task: knowledge_id=%s, file_path=%s", payload.KnowledgeID, payload.FilePath)
|
||||
|
||||
// 幂等性检查:获取knowledge记录
|
||||
knowledge, err := s.repo.GetKnowledgeByID(ctx, payload.TenantID, payload.KnowledgeID)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "failed to get knowledge: %v", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
if knowledge == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 检查任务状态 - 幂等性处理
|
||||
if knowledge.ParseStatus == "completed" {
|
||||
logger.Infof(ctx, "Document already completed, skipping: %s", payload.KnowledgeID)
|
||||
return nil // 幂等:已完成的任务直接返回
|
||||
}
|
||||
|
||||
if knowledge.ParseStatus == "failed" {
|
||||
// 检查是否可恢复(例如:超时、临时错误等)
|
||||
// 对于不可恢复的错误,直接返回
|
||||
logger.Warnf(ctx, "Document processing previously failed: %s, error: %s", payload.KnowledgeID, knowledge.ErrorMessage)
|
||||
// 这里可以根据错误类型判断是否可恢复,暂时允许重试
|
||||
}
|
||||
|
||||
// 检查是否有部分处理(有chunks但状态不是completed)
|
||||
if knowledge.ParseStatus != "completed" && knowledge.ParseStatus != "pending" && knowledge.ParseStatus != "processing" {
|
||||
// 状态异常,记录日志但继续处理
|
||||
logger.Warnf(ctx, "Unexpected parse status: %s for knowledge: %s", knowledge.ParseStatus, payload.KnowledgeID)
|
||||
}
|
||||
|
||||
// 获取知识库信息
|
||||
kb, err := s.kbService.GetKnowledgeBaseByID(ctx, payload.KnowledgeBaseID)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "failed to get knowledge base: %v", err)
|
||||
knowledge.ParseStatus = "failed"
|
||||
knowledge.ErrorMessage = fmt.Sprintf("failed to get knowledge base: %v", err)
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
s.repo.UpdateKnowledge(ctx, knowledge)
|
||||
return nil
|
||||
}
|
||||
|
||||
knowledge.ParseStatus = "processing"
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
if err := s.repo.UpdateKnowledge(ctx, knowledge); err != nil {
|
||||
logger.Errorf(ctx, "failed to update knowledge status to processing: %v", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
// 构建VLM配置(如果需要)
|
||||
var vlmConfig *proto.VLMConfig
|
||||
if payload.EnableMultimodel {
|
||||
vlmConfig, err = s.getVLMProtoConfig(ctx, kb)
|
||||
if err != nil {
|
||||
logger.GetLogger(ctx).WithField("knowledge_id", knowledge.ID).
|
||||
WithField("error", err).Errorf("processDocument build VLM config failed")
|
||||
knowledge.ParseStatus = "failed"
|
||||
knowledge.ErrorMessage = err.Error()
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
s.repo.UpdateKnowledge(ctx, knowledge)
|
||||
return nil
|
||||
}
|
||||
if vlmConfig == nil {
|
||||
logger.GetLogger(ctx).WithField("knowledge_id", knowledge.ID).
|
||||
Error("processDocument enable multimodal but VLM config missing")
|
||||
knowledge.ParseStatus = "failed"
|
||||
knowledge.ErrorMessage = "VLM 配置缺失"
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
s.repo.UpdateKnowledge(ctx, knowledge)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// 检查多模态配置(仅对文件导入)
|
||||
if payload.FilePath != "" && !payload.EnableMultimodel && IsImageType(payload.FileType) {
|
||||
logger.GetLogger(ctx).WithField("knowledge_id", knowledge.ID).
|
||||
WithField("error", ErrImageNotParse).Errorf("processDocument image without enable multimodel")
|
||||
knowledge.ParseStatus = "failed"
|
||||
knowledge.ErrorMessage = ErrImageNotParse.Error()
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
s.repo.UpdateKnowledge(ctx, knowledge)
|
||||
return nil
|
||||
}
|
||||
|
||||
// 处理不同类型的导入:文件、URL、文本段落
|
||||
var chunks []*proto.Chunk
|
||||
if payload.URL != "" {
|
||||
// URL导入
|
||||
urlResp, err := s.docReaderClient.ReadFromURL(ctx, &proto.ReadFromURLRequest{
|
||||
Url: payload.URL,
|
||||
Title: knowledge.Title,
|
||||
ReadConfig: &proto.ReadConfig{
|
||||
ChunkSize: int32(kb.ChunkingConfig.ChunkSize),
|
||||
ChunkOverlap: int32(kb.ChunkingConfig.ChunkOverlap),
|
||||
Separators: kb.ChunkingConfig.Separators,
|
||||
EnableMultimodal: payload.EnableMultimodel,
|
||||
StorageConfig: &proto.StorageConfig{
|
||||
Provider: proto.StorageProvider(proto.StorageProvider_value[strings.ToUpper(kb.StorageConfig.Provider)]),
|
||||
Region: kb.StorageConfig.Region,
|
||||
BucketName: kb.StorageConfig.BucketName,
|
||||
AccessKeyId: kb.StorageConfig.SecretID,
|
||||
SecretAccessKey: kb.StorageConfig.SecretKey,
|
||||
AppId: kb.StorageConfig.AppID,
|
||||
PathPrefix: kb.StorageConfig.PathPrefix,
|
||||
},
|
||||
VlmConfig: vlmConfig,
|
||||
},
|
||||
RequestId: payload.RequestId,
|
||||
})
|
||||
if err != nil {
|
||||
knowledge.ParseStatus = "failed"
|
||||
knowledge.ErrorMessage = err.Error()
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
s.repo.UpdateKnowledge(ctx, knowledge)
|
||||
return fmt.Errorf("failed to read from URL: %w", err)
|
||||
}
|
||||
chunks = urlResp.Chunks
|
||||
} else if len(payload.Passages) > 0 {
|
||||
// 文本段落导入
|
||||
chunks := make([]*proto.Chunk, 0, len(payload.Passages))
|
||||
start, end := 0, 0
|
||||
for i, p := range payload.Passages {
|
||||
if p == "" {
|
||||
continue
|
||||
}
|
||||
end += len([]rune(p))
|
||||
chunk := &proto.Chunk{
|
||||
Content: p,
|
||||
Seq: int32(i),
|
||||
Start: int32(start),
|
||||
End: int32(end),
|
||||
}
|
||||
start = end
|
||||
chunks = append(chunks, chunk)
|
||||
}
|
||||
// 直接处理chunks,不需要调用docReader
|
||||
s.processChunks(ctx, kb, knowledge, chunks)
|
||||
return nil
|
||||
} else {
|
||||
// 文件导入
|
||||
fileReader, err := s.fileSvc.GetFile(ctx, payload.FilePath)
|
||||
if err != nil {
|
||||
logger.GetLogger(ctx).WithField("knowledge_id", knowledge.ID).
|
||||
WithField("error", err).Errorf("processDocument get file failed")
|
||||
knowledge.ParseStatus = "failed"
|
||||
knowledge.ErrorMessage = err.Error()
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
s.repo.UpdateKnowledge(ctx, knowledge)
|
||||
return fmt.Errorf("failed to get file: %w", err)
|
||||
}
|
||||
defer fileReader.Close()
|
||||
|
||||
// 读取文件内容
|
||||
contentBytes, err := io.ReadAll(fileReader)
|
||||
if err != nil {
|
||||
knowledge.ParseStatus = "failed"
|
||||
knowledge.ErrorMessage = err.Error()
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
s.repo.UpdateKnowledge(ctx, knowledge)
|
||||
return fmt.Errorf("failed to read file: %w", err)
|
||||
}
|
||||
|
||||
// 调用docReader处理文件
|
||||
fileResp, err := s.docReaderClient.ReadFromFile(ctx, &proto.ReadFromFileRequest{
|
||||
FileContent: contentBytes,
|
||||
FileName: payload.FileName,
|
||||
FileType: payload.FileType,
|
||||
ReadConfig: &proto.ReadConfig{
|
||||
ChunkSize: int32(kb.ChunkingConfig.ChunkSize),
|
||||
ChunkOverlap: int32(kb.ChunkingConfig.ChunkOverlap),
|
||||
Separators: kb.ChunkingConfig.Separators,
|
||||
EnableMultimodal: payload.EnableMultimodel,
|
||||
StorageConfig: &proto.StorageConfig{
|
||||
Provider: proto.StorageProvider(proto.StorageProvider_value[strings.ToUpper(kb.StorageConfig.Provider)]),
|
||||
Region: kb.StorageConfig.Region,
|
||||
BucketName: kb.StorageConfig.BucketName,
|
||||
AccessKeyId: kb.StorageConfig.SecretID,
|
||||
SecretAccessKey: kb.StorageConfig.SecretKey,
|
||||
AppId: kb.StorageConfig.AppID,
|
||||
PathPrefix: kb.StorageConfig.PathPrefix,
|
||||
},
|
||||
VlmConfig: vlmConfig,
|
||||
},
|
||||
RequestId: payload.RequestId,
|
||||
})
|
||||
if err != nil {
|
||||
logger.GetLogger(ctx).WithField("knowledge_id", knowledge.ID).
|
||||
WithField("error", err).Errorf("processDocument read file failed")
|
||||
knowledge.ParseStatus = "failed"
|
||||
knowledge.ErrorMessage = err.Error()
|
||||
knowledge.UpdatedAt = time.Now()
|
||||
s.repo.UpdateKnowledge(ctx, knowledge)
|
||||
return fmt.Errorf("failed to read file from docreader: %w", err)
|
||||
}
|
||||
chunks = fileResp.Chunks
|
||||
}
|
||||
|
||||
// 处理chunks(这会更新状态为completed)
|
||||
s.processChunks(ctx, kb, knowledge, chunks)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ProcessFAQImport handles Asynq FAQ import tasks
|
||||
func (s *knowledgeService) ProcessFAQImport(ctx context.Context, t *asynq.Task) error {
|
||||
var payload types.FAQImportPayload
|
||||
if err := json.Unmarshal(t.Payload(), &payload); err != nil {
|
||||
logger.Errorf(ctx, "failed to unmarshal FAQ import task payload: %v", err)
|
||||
return fmt.Errorf("failed to unmarshal task payload: %w", err)
|
||||
}
|
||||
|
||||
ctx = logger.WithRequestID(ctx, uuid.New().String())
|
||||
ctx = logger.WithField(ctx, "faq_import", payload.TaskID)
|
||||
ctx = context.WithValue(ctx, types.TenantIDContextKey, payload.TenantID)
|
||||
|
||||
tenantInfo, err := s.tenantRepo.GetTenantByID(ctx, payload.TenantID)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "failed to get tenant: %v", err)
|
||||
return nil
|
||||
}
|
||||
ctx = context.WithValue(ctx, types.TenantInfoContextKey, tenantInfo)
|
||||
|
||||
logger.Infof(ctx, "Processing FAQ import task: task_id=%s, kb_id=%s, total_entries=%d", payload.TaskID, payload.KBID, len(payload.Entries))
|
||||
|
||||
// 幂等性检查:获取knowledge记录(FAQ任务使用knowledge ID作为taskID)
|
||||
knowledge, err := s.repo.GetKnowledgeByID(ctx, payload.TenantID, payload.TaskID)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "failed to get FAQ knowledge: %v", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
if knowledge == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 检查任务状态 - 幂等性处理
|
||||
if knowledge.ParseStatus == "completed" {
|
||||
logger.Infof(ctx, "FAQ import already completed, skipping: %s", payload.TaskID)
|
||||
return nil // 幂等:已完成的任务直接返回
|
||||
}
|
||||
|
||||
// 检查已处理进度
|
||||
importMeta, _ := types.ParseFAQImportMetadata(knowledge)
|
||||
var processedCount int
|
||||
if importMeta != nil {
|
||||
processedCount = importMeta.ImportProcessed
|
||||
logger.Infof(ctx, "Resuming FAQ import from progress: %d/%d", processedCount, len(payload.Entries))
|
||||
}
|
||||
|
||||
// 保存原始总数量(在截断payload.Entries之前)
|
||||
originalTotalEntries := len(payload.Entries)
|
||||
|
||||
// 如果已经处理了一部分,需要从该位置继续
|
||||
if processedCount < originalTotalEntries {
|
||||
// 幂等性处理:清理processedCount之后可能已部分处理的chunks和索引数据
|
||||
// 因为任务可能在处理过程中中断,导致部分chunks已入库或已写入索引
|
||||
logger.Infof(ctx, "Cleaning up potentially partially processed chunks after entry %d", processedCount)
|
||||
|
||||
// 计算需要清理的ChunkIndex范围
|
||||
// 已处理的chunks的ChunkIndex范围是 [payload.StartChunkIndex, payload.StartChunkIndex + processedCount - 1]
|
||||
// 需要清理的chunks是从已处理位置之后开始,到分配的结束位置
|
||||
minChunkIndex := payload.StartChunkIndex + processedCount
|
||||
maxChunkIndex := payload.EndChunkIndex
|
||||
logger.Infof(ctx, "Cleaning chunks with ChunkIndex in range [%d, %d] (allocated range: [%d, %d], processedCount=%d), knowledge_id=%s, kb_id=%s",
|
||||
minChunkIndex, maxChunkIndex, payload.StartChunkIndex, payload.EndChunkIndex, processedCount, payload.KnowledgeID, payload.KBID)
|
||||
|
||||
// 查询并删除需要清理的chunks(使用ChunkIndex范围查询)
|
||||
// 注意:只清理processedCount之后的部分,避免删除已成功处理的chunks
|
||||
chunksToDelete, err := s.chunkRepo.DeleteChunksByChunkIndexRange(ctx,
|
||||
payload.TenantID,
|
||||
payload.KnowledgeID,
|
||||
minChunkIndex,
|
||||
maxChunkIndex,
|
||||
)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to query chunks for cleanup: %v", err)
|
||||
return fmt.Errorf("failed to query chunks: %w", err)
|
||||
}
|
||||
logger.Infof(ctx, "Cleaned %d chunks, from knowledge: %s, range: [%d, %d]",
|
||||
len(chunksToDelete), payload.KnowledgeID, minChunkIndex, maxChunkIndex)
|
||||
|
||||
// 获取KB信息以删除索引数据
|
||||
kb, err := s.kbService.GetKnowledgeBaseByID(ctx, payload.KBID)
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to get knowledge base: %v", err)
|
||||
return fmt.Errorf("failed to get knowledge base: %w", err)
|
||||
}
|
||||
|
||||
// 删除索引数据
|
||||
embeddingModel, err := s.modelService.GetEmbeddingModel(ctx, kb.EmbeddingModelID)
|
||||
if err == nil {
|
||||
retrieveEngine, err := retriever.NewCompositeRetrieveEngine(s.retrieveEngine, tenantInfo.RetrieverEngines.Engines)
|
||||
if err == nil {
|
||||
chunkIDs := make([]string, 0, len(chunksToDelete))
|
||||
for _, chunk := range chunksToDelete {
|
||||
chunkIDs = append(chunkIDs, chunk.ID)
|
||||
}
|
||||
if err := retrieveEngine.DeleteByChunkIDList(ctx, chunkIDs, embeddingModel.GetDimensions()); err != nil {
|
||||
logger.Warnf(ctx, "Failed to delete index data for chunks (may not exist): %v", err)
|
||||
} else {
|
||||
logger.Infof(ctx, "Successfully deleted index data for %d chunks", len(chunksToDelete))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 从已处理的位置继续
|
||||
payload.Entries = payload.Entries[processedCount:]
|
||||
logger.Infof(ctx, "Continuing FAQ import from entry %d, remaining: %d entries", processedCount, len(payload.Entries))
|
||||
}
|
||||
|
||||
// 更新任务状态为运行中
|
||||
if err := s.updateFAQImportStatusWithRanges(ctx, payload.TaskID, types.FAQImportStatusProcessing, 0,
|
||||
originalTotalEntries, processedCount, "", importMeta.NextChunkIndex); err != nil {
|
||||
logger.Errorf(ctx, "Failed to update task status to running: %v", err)
|
||||
}
|
||||
|
||||
// 构建FAQBatchUpsertPayload
|
||||
faqPayload := &types.FAQBatchUpsertPayload{
|
||||
Entries: payload.Entries,
|
||||
Mode: payload.Mode,
|
||||
}
|
||||
|
||||
// 执行FAQ导入(传入原始总数量)
|
||||
if err := s.executeFAQImport(ctx, payload.TaskID, payload.KBID, faqPayload, payload.TenantID, payload.StartChunkIndex, processedCount, originalTotalEntries); err != nil {
|
||||
logger.Errorf(ctx, "FAQ import task failed: %s, error: %v", payload.TaskID, err)
|
||||
return fmt.Errorf("FAQ import failed: %w", err)
|
||||
}
|
||||
|
||||
// 任务成功完成
|
||||
logger.Infof(ctx, "FAQ import task completed: %s", payload.TaskID)
|
||||
if err := s.updateFAQImportStatus(ctx, payload.TaskID, types.FAQImportStatusCompleted, 100, originalTotalEntries, originalTotalEntries, ""); err != nil {
|
||||
logger.Errorf(ctx, "Failed to update task status to success: %v", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -454,6 +454,8 @@ func (s *knowledgeBaseService) HybridSearch(ctx context.Context,
|
||||
return nil, err
|
||||
}
|
||||
|
||||
matchCount := params.MatchCount * 3
|
||||
|
||||
// Add vector retrieval params if supported
|
||||
if retrieveEngine.SupportRetriever(types.VectorRetrieverType) && !params.DisableVectorMatch {
|
||||
logger.Info(ctx, "Vector retrieval supported, preparing vector retrieval parameters")
|
||||
@@ -479,7 +481,7 @@ func (s *knowledgeBaseService) HybridSearch(ctx context.Context,
|
||||
Query: params.QueryText,
|
||||
Embedding: queryEmbedding,
|
||||
KnowledgeBaseIDs: []string{id},
|
||||
TopK: params.MatchCount,
|
||||
TopK: matchCount,
|
||||
Threshold: params.VectorThreshold,
|
||||
RetrieverType: types.VectorRetrieverType,
|
||||
})
|
||||
@@ -492,7 +494,7 @@ func (s *knowledgeBaseService) HybridSearch(ctx context.Context,
|
||||
retrieveParams = append(retrieveParams, types.RetrieveParams{
|
||||
Query: params.QueryText,
|
||||
KnowledgeBaseIDs: []string{id},
|
||||
TopK: params.MatchCount,
|
||||
TopK: matchCount,
|
||||
Threshold: params.KeywordThreshold,
|
||||
RetrieverType: types.KeywordsRetrieverType,
|
||||
})
|
||||
@@ -543,7 +545,7 @@ func (s *knowledgeBaseService) HybridSearch(ctx context.Context,
|
||||
// Check if we need iterative retrieval for FAQ with separate indexing
|
||||
// Only use iterative retrieval if we don't have enough unique chunks after first deduplication
|
||||
needsIterativeRetrieval := len(deduplicatedChunks) < params.MatchCount &&
|
||||
kb.Type == types.KnowledgeBaseTypeFAQ && len(matchResults) >= params.MatchCount
|
||||
kb.Type == types.KnowledgeBaseTypeFAQ && len(matchResults) == matchCount
|
||||
if needsIterativeRetrieval {
|
||||
logger.Infof(ctx, "Not enough unique chunks (%d < %d), using iterative retrieval for FAQ",
|
||||
len(deduplicatedChunks), params.MatchCount)
|
||||
|
||||
@@ -14,8 +14,9 @@ import (
|
||||
type AsynqTaskParams struct {
|
||||
dig.In
|
||||
|
||||
Server *asynq.Server
|
||||
Extracter interfaces.Extracter
|
||||
Server *asynq.Server
|
||||
Extracter interfaces.Extracter
|
||||
KnowledgeService interfaces.KnowledgeService
|
||||
}
|
||||
|
||||
func getAsynqRedisClientOpt() *asynq.RedisClientOpt {
|
||||
@@ -60,6 +61,12 @@ func RunAsynqServer(params AsynqTaskParams) *asynq.ServeMux {
|
||||
|
||||
mux.HandleFunc(types.TypeChunkExtract, params.Extracter.Extract)
|
||||
|
||||
// Register document processing handler
|
||||
mux.HandleFunc(types.TypeDocumentProcess, params.KnowledgeService.ProcessDocument)
|
||||
|
||||
// Register FAQ import handler
|
||||
mux.HandleFunc(types.TypeFAQImport, params.KnowledgeService.ProcessFAQImport)
|
||||
|
||||
go func() {
|
||||
// Start the server
|
||||
if err := params.Server.Run(mux); err != nil {
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
package types
|
||||
|
||||
const (
|
||||
TypeChunkExtract = "chunk:extract"
|
||||
TypeChunkExtract = "chunk:extract"
|
||||
TypeDocumentProcess = "document:process" // 文档处理任务
|
||||
TypeFAQImport = "faq:import" // FAQ导入任务
|
||||
)
|
||||
|
||||
type ExtractChunkPayload struct {
|
||||
@@ -10,6 +12,32 @@ type ExtractChunkPayload struct {
|
||||
ModelID string `json:"model_id"`
|
||||
}
|
||||
|
||||
// DocumentProcessPayload 文档处理任务payload
|
||||
type DocumentProcessPayload struct {
|
||||
RequestId string `json:"request_id"`
|
||||
TenantID uint `json:"tenant_id"`
|
||||
KnowledgeID string `json:"knowledge_id"`
|
||||
KnowledgeBaseID string `json:"knowledge_base_id"`
|
||||
FilePath string `json:"file_path,omitempty"` // 文件路径(文件导入时使用)
|
||||
FileName string `json:"file_name,omitempty"` // 文件名(文件导入时使用)
|
||||
FileType string `json:"file_type,omitempty"` // 文件类型(文件导入时使用)
|
||||
URL string `json:"url,omitempty"` // URL(URL导入时使用)
|
||||
Passages []string `json:"passages,omitempty"` // 文本段落(文本导入时使用)
|
||||
EnableMultimodel bool `json:"enable_multimodel"`
|
||||
}
|
||||
|
||||
// FAQImportPayload FAQ导入任务payload
|
||||
type FAQImportPayload struct {
|
||||
TenantID uint `json:"tenant_id"`
|
||||
TaskID string `json:"task_id"`
|
||||
KBID string `json:"kb_id"`
|
||||
KnowledgeID string `json:"knowledge_id"`
|
||||
Entries []FAQEntryPayload `json:"entries"`
|
||||
Mode string `json:"mode"`
|
||||
StartChunkIndex int `json:"start_chunk_index"`
|
||||
EndChunkIndex int `json:"end_chunk_index"`
|
||||
}
|
||||
|
||||
type PromptTemplateStructured struct {
|
||||
Description string `json:"description"`
|
||||
Tags []string `json:"tags"`
|
||||
|
||||
@@ -114,10 +114,10 @@ type FAQSearchRequest struct {
|
||||
type FAQImportTaskStatus string
|
||||
|
||||
const (
|
||||
FAQImportStatusPending FAQImportTaskStatus = "pending"
|
||||
FAQImportStatusRunning FAQImportTaskStatus = "running"
|
||||
FAQImportStatusSuccess FAQImportTaskStatus = "success"
|
||||
FAQImportStatusFailed FAQImportTaskStatus = "failed"
|
||||
FAQImportStatusPending FAQImportTaskStatus = "pending"
|
||||
FAQImportStatusProcessing FAQImportTaskStatus = "processing"
|
||||
FAQImportStatusCompleted FAQImportTaskStatus = "completed"
|
||||
FAQImportStatusFailed FAQImportTaskStatus = "failed"
|
||||
)
|
||||
|
||||
// FAQImportMetadata 存储在Knowledge.Metadata中的FAQ导入任务信息
|
||||
@@ -125,6 +125,7 @@ type FAQImportMetadata struct {
|
||||
ImportProgress int `json:"import_progress"` // 0-100
|
||||
ImportTotal int `json:"import_total"`
|
||||
ImportProcessed int `json:"import_processed"`
|
||||
NextChunkIndex int `json:"next_chunk_index,omitempty"` // 下一个要分配的ChunkIndex(自增计数器)
|
||||
}
|
||||
|
||||
// ToJSON converts the metadata to JSON type.
|
||||
|
||||
@@ -41,6 +41,7 @@ type ChunkRepository interface {
|
||||
DeleteByKnowledgeList(ctx context.Context, tenantID uint, knowledgeIDs []string) error
|
||||
// CountChunksByKnowledgeBaseID counts the number of chunks in a knowledge base.
|
||||
CountChunksByKnowledgeBaseID(ctx context.Context, tenantID uint, kbID string) (int64, error)
|
||||
DeleteChunksByChunkIndexRange(ctx context.Context, tenantID uint, knowledgeID string, startChunkIndex int, endChunkIndex int) ([]*types.Chunk, error)
|
||||
}
|
||||
|
||||
// ChunkService defines the interface for chunk service operations
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"mime/multipart"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/hibiken/asynq"
|
||||
)
|
||||
|
||||
// KnowledgeService defines the interface for knowledge services.
|
||||
@@ -73,6 +74,10 @@ type KnowledgeService interface {
|
||||
UpdateFAQEntryTagBatch(ctx context.Context, kbID string, updates map[string]*string) error
|
||||
// GetRepository gets the knowledge repository
|
||||
GetRepository() KnowledgeRepository
|
||||
// ProcessDocument handles Asynq document processing tasks
|
||||
ProcessDocument(ctx context.Context, t *asynq.Task) error
|
||||
// ProcessFAQImport handles Asynq FAQ import tasks
|
||||
ProcessFAQImport(ctx context.Context, t *asynq.Task) error
|
||||
}
|
||||
|
||||
// KnowledgeRepository defines the interface for knowledge repositories.
|
||||
|
||||
Reference in New Issue
Block a user