Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 21 additions & 4 deletions internal/handler/knowledge.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 扫描知识库
Expand Down
95 changes: 78 additions & 17 deletions internal/knowledge/indexer.go
Original file line number Diff line number Diff line change
Expand Up @@ -213,29 +213,53 @@ 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
idx.rebuildFailed = 0
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()
Expand All @@ -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
Expand Down Expand Up @@ -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
}

Expand Down
20 changes: 20 additions & 0 deletions internal/knowledge/indexer_rebuild_state_test.go
Original file line number Diff line number Diff line change
@@ -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()
}
29 changes: 19 additions & 10 deletions internal/security/executor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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" // 默认格式
Expand All @@ -418,38 +430,35 @@ 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 {
// 如果没有模板,使用默认格式
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)
}
}

Expand Down