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:
wizardchen
2025-11-24 00:21:42 +08:00
parent 057cfd9eb0
commit 4a2be2b878
12 changed files with 790 additions and 532 deletions
+2 -2
View File
@@ -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 }) {
+53 -17
View File
@@ -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) {
+2 -1
View File
@@ -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,
});
}
+77 -329
View File
@@ -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;
+19
View File
@@ -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
}
+583 -173
View File
@@ -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)
+9 -2
View File
@@ -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 {
+29 -1
View File
@@ -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"`
+5 -4
View File
@@ -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.
+1
View File
@@ -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
+5
View File
@@ -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.