feat: add knowledge_delete tool + cmd_run self-kill guard via before_toolcall stage hook

This commit is contained in:
root
2026-07-16 19:24:31 +08:00
parent c28aac2519
commit 4184cd0849
3 changed files with 62 additions and 20 deletions

View File

@ -1464,6 +1464,16 @@ func (a *Agent) executeKnowledgeTool(tc agentAPI.ToolCall) string {
tree := a.knowledge.BuildTree() tree := a.knowledge.BuildTree()
return formatTree(tree, 0) return formatTree(tree, 0)
case "knowledge_delete":
name, _ := tc.Arguments["name"].(string)
if name == "" {
return "name 不能为空"
}
if err := a.knowledge.Remove(name); err != nil {
return fmt.Sprintf("知识删除失败: %v", err)
}
return fmt.Sprintf("知识「%s」已删除", name)
default: default:
return fmt.Sprintf("未知的知识工具: %s", tc.Name) return fmt.Sprintf("未知的知识工具: %s", tc.Name)
} }
@ -1833,6 +1843,20 @@ func (a *Agent) buildToolDefs() []interface{} {
}, },
}, },
}) })
tools = append(tools, map[string]interface{}{
"type": "function",
"function": map[string]interface{}{
"name": "knowledge_delete",
"description": "删除知识库中的指定知识条目。",
"parameters": map[string]interface{}{
"type": "object",
"properties": map[string]interface{}{
"name": map[string]interface{}{"type": "string", "description": "要删除的知识名称"},
},
"required": []string{"name"},
},
},
})
} }
// 文档记忆工具 // 文档记忆工具

View File

@ -224,9 +224,11 @@ func (s *Store) Remove(name string) error {
} }
delete(s.items, id) delete(s.items, id)
s.vec.Remove(id) s.vec.Remove(id)
if err := s.writeIndex(); err != nil { go func() {
log.Printf("[knowledge] write index error after removing %s: %v", name, err) if err := s.writeIndex(); err != nil {
} log.Printf("[knowledge] write index error after removing %s: %v", name, err)
}
}()
return nil return nil
} }

View File

@ -134,10 +134,6 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
return map[string]interface{}{"error": "command is required"}, nil return map[string]interface{}{"error": "command is required"}, nil
} }
if blocked, reason := isDangerousCommand(command); blocked {
return map[string]interface{}{"error": reason, "status": "denied"}, nil
}
timeoutStr, _ := args["timeout"].(string) timeoutStr, _ := args["timeout"].(string)
if timeoutStr == "" { if timeoutStr == "" {
timeoutStr = p.defaultTimeout timeoutStr = p.defaultTimeout
@ -218,6 +214,27 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error {
}, nil }, nil
}) })
s.RegisterStageOwnTools(sdk.StageBeforeToolcall, func(ctx *sdk.StageContext) error {
if len(ctx.ToolCalls) == 0 {
return nil
}
tc := ctx.ToolCalls[0]
if tc.Name != "cmd_run" {
return nil
}
command, _ := tc.Arguments["command"].(string)
if command == "" {
return nil
}
if blocked, reason := isDangerousCommand(command); blocked {
ctx.Lock()
ctx.Response = &reason
ctx.Unlock()
return nil
}
return nil
})
return nil return nil
} }
@ -250,9 +267,9 @@ func isDangerousCommand(command string) (bool, string) {
pidStr := fmt.Sprintf("%d", pid) pidStr := fmt.Sprintf("%d", pid)
lower := strings.ToLower(command) lower := strings.ToLower(command)
dangerousPatterns := []string{ selfTargets := []string{"homed", pidStr}
"homed", killCmds := []string{"kill ", "killall ", "pkill ", "kill -", "kill -9"}
pidStr, serviceCmds := []string{
"systemctl stop homeagent", "systemctl stop homeagent",
"systemctl restart homeagent", "systemctl restart homeagent",
"systemctl kill homeagent", "systemctl kill homeagent",
@ -260,20 +277,19 @@ func isDangerousCommand(command string) (bool, string) {
"service homeagent restart", "service homeagent restart",
} }
killPatterns := []string{"kill ", "killall ", "pkill ", "kill -", "kill -9"} for _, kp := range killCmds {
if !strings.Contains(lower, kp) {
for _, kp := range killPatterns { continue
if strings.Contains(lower, kp) { }
for _, target := range dangerousPatterns { for _, t := range selfTargets {
if strings.Contains(lower, strings.ToLower(target)) { if strings.Contains(lower, strings.ToLower(t)) {
return true, fmt.Sprintf("命令被拦截:禁止向 homed 进程(pid=%s)发送信号", pidStr) return true, fmt.Sprintf("命令被拦截:禁止向 homed 进程(pid=%s)发送信号", pidStr)
}
} }
} }
} }
for _, dp := range dangerousPatterns[2:] { for _, sc := range serviceCmds {
if strings.Contains(lower, dp) { if strings.Contains(lower, sc) {
return true, fmt.Sprintf("命令被拦截:禁止操作 homed 服务进程(pid=%s)", pidStr) return true, fmt.Sprintf("命令被拦截:禁止操作 homed 服务进程(pid=%s)", pidStr)
} }
} }