From 7a44ea168993252847fe9f79aa8fbba60a836dc8 Mon Sep 17 00:00:00 2001 From: ddick2000 Date: Wed, 8 Jul 2026 09:24:09 +0800 Subject: [PATCH] fix resumable knowledge indexing and empty flags --- internal/handler/knowledge.go | 25 ++++- internal/knowledge/indexer.go | 95 +++++++++++++++---- .../knowledge/indexer_rebuild_state_test.go | 20 ++++ internal/security/executor.go | 29 ++++-- 4 files changed, 138 insertions(+), 31 deletions(-) create mode 100644 internal/knowledge/indexer_rebuild_state_test.go diff --git a/internal/handler/knowledge.go b/internal/handler/knowledge.go index eee106ac6..f9315b0c4 100644 --- a/internal/handler/knowledge.go +++ b/internal/handler/knowledge.go @@ -318,18 +318,35 @@ func (h *KnowledgeHandler) DeleteItem(c *gin.Context) { // RebuildIndex 重建索引 func (h *KnowledgeHandler) RebuildIndex(c *gin.Context) { - // 异步重建索引 + if isRebuilding, _, _, _, _, _, _ := h.indexer.GetRebuildStatus(); isRebuilding { + c.JSON(http.StatusConflict, gin.H{"error": "已有索引任务正在进行,请等待完成"}) + return + } + + mode := c.Query("mode") + fullRebuild := mode == "full" || c.Query("full") == "true" + + message := "缺失索引补齐已开始,将在后台进行" go func() { ctx := context.Background() - if err := h.indexer.RebuildIndex(ctx); err != nil { - h.logger.Error("重建索引失败", zap.Error(err)) + if fullRebuild { + if err := h.indexer.RebuildIndex(ctx); err != nil { + h.logger.Error("重建索引失败", zap.Error(err)) + } + return + } + if err := h.indexer.IndexMissing(ctx); err != nil { + h.logger.Error("补齐缺失索引失败", zap.Error(err)) } }() + if fullRebuild { + message = "全量索引重建已开始,将在后台进行" + } if h.audit != nil { h.audit.RecordOK(c, "knowledge", "index_rebuild", "重建知识库索引", "knowledge", "", nil) } - c.JSON(http.StatusOK, gin.H{"message": "索引重建已开始,将在后台进行"}) + c.JSON(http.StatusOK, gin.H{"message": message, "mode": mode}) } // ScanKnowledgeBase 扫描知识库 diff --git a/internal/knowledge/indexer.go b/internal/knowledge/indexer.go index aeb6b9ff0..61f4f1f1f 100644 --- a/internal/knowledge/indexer.go +++ b/internal/knowledge/indexer.go @@ -213,9 +213,13 @@ func (idx *Indexer) HasIndex() (bool, error) { return count > 0, nil } -// RebuildIndex 重建所有索引 -func (idx *Indexer) RebuildIndex(ctx context.Context) error { +func (idx *Indexer) beginIndexRun() error { idx.rebuildMu.Lock() + defer idx.rebuildMu.Unlock() + + if idx.isRebuilding { + return fmt.Errorf("索引任务已在进行中") + } idx.isRebuilding = true idx.rebuildTotalItems = 0 idx.rebuildCurrent = 0 @@ -223,19 +227,39 @@ func (idx *Indexer) RebuildIndex(ctx context.Context) error { idx.rebuildStartTime = time.Now() idx.rebuildLastItemID = "" idx.rebuildLastChunks = 0 + return nil +} + +func (idx *Indexer) finishIndexRun() { + idx.rebuildMu.Lock() + idx.isRebuilding = false idx.rebuildMu.Unlock() +} +func (idx *Indexer) resetLastError() { idx.mu.Lock() idx.lastError = "" idx.lastErrorTime = time.Time{} idx.errorCount = 0 idx.mu.Unlock() +} - rows, err := idx.db.Query("SELECT id FROM knowledge_base_items") +func (idx *Indexer) setIndexRunTotal(total int) { + idx.rebuildMu.Lock() + idx.rebuildTotalItems = total + idx.rebuildMu.Unlock() +} + +// RebuildIndex 重建所有索引 +func (idx *Indexer) RebuildIndex(ctx context.Context) error { + if err := idx.beginIndexRun(); err != nil { + return err + } + defer idx.finishIndexRun() + idx.resetLastError() + + rows, err := idx.db.QueryContext(ctx, "SELECT id FROM knowledge_base_items ORDER BY updated_at ASC, id ASC") if err != nil { - idx.rebuildMu.Lock() - idx.isRebuilding = false - idx.rebuildMu.Unlock() return fmt.Errorf("查询知识项失败:%w", err) } defer rows.Close() @@ -244,20 +268,61 @@ func (idx *Indexer) RebuildIndex(ctx context.Context) error { for rows.Next() { var id string if err := rows.Scan(&id); err != nil { - idx.rebuildMu.Lock() - idx.isRebuilding = false - idx.rebuildMu.Unlock() return fmt.Errorf("扫描知识项 ID 失败:%w", err) } itemIDs = append(itemIDs, id) } + if err := rows.Err(); err != nil { + return fmt.Errorf("扫描知识项 ID 失败:%w", err) + } - idx.rebuildMu.Lock() - idx.rebuildTotalItems = len(itemIDs) - idx.rebuildMu.Unlock() + idx.setIndexRunTotal(len(itemIDs)) idx.logger.Info("开始重建索引", zap.Int("totalItems", len(itemIDs))) + return idx.indexItemIDs(ctx, itemIDs, "索引重建完成") +} + +// IndexMissing 只为还没有向量的知识项补齐索引,适合中断后续跑。 +func (idx *Indexer) IndexMissing(ctx context.Context) error { + if err := idx.beginIndexRun(); err != nil { + return err + } + defer idx.finishIndexRun() + idx.resetLastError() + + rows, err := idx.db.QueryContext(ctx, ` + SELECT i.id + FROM knowledge_base_items i + LEFT JOIN knowledge_embeddings e ON e.item_id = i.id + WHERE e.item_id IS NULL + ORDER BY i.updated_at ASC, i.id ASC + `) + if err != nil { + return fmt.Errorf("查询未索引知识项失败:%w", err) + } + defer rows.Close() + + var itemIDs []string + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return fmt.Errorf("扫描未索引知识项 ID 失败:%w", err) + } + itemIDs = append(itemIDs, id) + } + if err := rows.Err(); err != nil { + return fmt.Errorf("扫描未索引知识项 ID 失败:%w", err) + } + + idx.setIndexRunTotal(len(itemIDs)) + + idx.logger.Info("开始补齐缺失索引", zap.Int("totalItems", len(itemIDs))) + + return idx.indexItemIDs(ctx, itemIDs, "缺失索引补齐完成") +} + +func (idx *Indexer) indexItemIDs(ctx context.Context, itemIDs []string, doneMessage string) error { failedCount := 0 consecutiveFailures := 0 maxConsecutiveFailures := 5 @@ -329,11 +394,7 @@ func (idx *Indexer) RebuildIndex(ctx context.Context) error { } } - idx.rebuildMu.Lock() - idx.isRebuilding = false - idx.rebuildMu.Unlock() - - idx.logger.Info("索引重建完成", zap.Int("totalItems", len(itemIDs)), zap.Int("failedCount", failedCount)) + idx.logger.Info(doneMessage, zap.Int("totalItems", len(itemIDs)), zap.Int("failedCount", failedCount)) return nil } diff --git a/internal/knowledge/indexer_rebuild_state_test.go b/internal/knowledge/indexer_rebuild_state_test.go new file mode 100644 index 000000000..2b08efaca --- /dev/null +++ b/internal/knowledge/indexer_rebuild_state_test.go @@ -0,0 +1,20 @@ +package knowledge + +import "testing" + +func TestIndexerRejectsConcurrentIndexRuns(t *testing.T) { + idx := &Indexer{} + + if err := idx.beginIndexRun(); err != nil { + t.Fatalf("first index run should start: %v", err) + } + if err := idx.beginIndexRun(); err == nil { + t.Fatal("second index run should be rejected while one is active") + } + + idx.finishIndexRun() + if err := idx.beginIndexRun(); err != nil { + t.Fatalf("index run should start again after finish: %v", err) + } + idx.finishIndexRun() +} diff --git a/internal/security/executor.go b/internal/security/executor.go index fa8c157b7..35b95b3a6 100644 --- a/internal/security/executor.go +++ b/internal/security/executor.go @@ -407,6 +407,18 @@ func (e *Executor) buildCommandArgs(toolName string, toolConfig *config.ToolConf } } + formattedValue := e.formatParamValue(param, value) + if strings.TrimSpace(formattedValue) == "" { + if param.Required { + e.logger.Warn("必需参数为空", + zap.String("tool", toolName), + zap.String("param", param.Name), + ) + return []string{} + } + continue + } + format := param.Format if format == "" { format = "flag" // 默认格式 @@ -418,23 +430,20 @@ func (e *Executor) buildCommandArgs(toolName string, toolConfig *config.ToolConf if param.Flag != "" { cmdArgs = append(cmdArgs, param.Flag) } - formattedValue := e.formatParamValue(param, value) - if formattedValue != "" { - cmdArgs = append(cmdArgs, formattedValue) - } + cmdArgs = append(cmdArgs, formattedValue) case "combined": // --flag=value 或 -f=value if param.Flag != "" { - cmdArgs = append(cmdArgs, fmt.Sprintf("%s=%s", param.Flag, e.formatParamValue(param, value))) + cmdArgs = append(cmdArgs, fmt.Sprintf("%s=%s", param.Flag, formattedValue)) } else { - cmdArgs = append(cmdArgs, e.formatParamValue(param, value)) + cmdArgs = append(cmdArgs, formattedValue) } case "template": // 使用模板字符串 if param.Template != "" { template := param.Template template = strings.ReplaceAll(template, "{flag}", param.Flag) - template = strings.ReplaceAll(template, "{value}", e.formatParamValue(param, value)) + template = strings.ReplaceAll(template, "{value}", formattedValue) template = strings.ReplaceAll(template, "{name}", param.Name) cmdArgs = append(cmdArgs, strings.Fields(template)...) } else { @@ -442,14 +451,14 @@ func (e *Executor) buildCommandArgs(toolName string, toolConfig *config.ToolConf if param.Flag != "" { cmdArgs = append(cmdArgs, param.Flag) } - cmdArgs = append(cmdArgs, e.formatParamValue(param, value)) + cmdArgs = append(cmdArgs, formattedValue) } case "positional": // 位置参数(已在上面处理) - cmdArgs = append(cmdArgs, e.formatParamValue(param, value)) + cmdArgs = append(cmdArgs, formattedValue) default: // 默认:直接添加值 - cmdArgs = append(cmdArgs, e.formatParamValue(param, value)) + cmdArgs = append(cmdArgs, formattedValue) } }