From 09b79811a0c0f5f1c19f9b13792b62e267e54265 Mon Sep 17 00:00:00 2001 From: root Date: Thu, 16 Jul 2026 19:24:31 +0800 Subject: [PATCH] feat: add knowledge_delete tool + cmd_run self-kill guard via before_toolcall stage hook --- internal/agent/core/agent.go | 24 ++++++++++++++++ internal/knowledge/knowledge.go | 8 ++++-- internal/plugins/cmd/plugin.go | 50 ++++++++++++++++++++++----------- 3 files changed, 62 insertions(+), 20 deletions(-) diff --git a/internal/agent/core/agent.go b/internal/agent/core/agent.go index 9aa659a..6f7a253 100644 --- a/internal/agent/core/agent.go +++ b/internal/agent/core/agent.go @@ -1464,6 +1464,16 @@ func (a *Agent) executeKnowledgeTool(tc agentAPI.ToolCall) string { tree := a.knowledge.BuildTree() 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: 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"}, + }, + }, + }) } // 文档记忆工具 diff --git a/internal/knowledge/knowledge.go b/internal/knowledge/knowledge.go index 444b672..f8c3d2a 100644 --- a/internal/knowledge/knowledge.go +++ b/internal/knowledge/knowledge.go @@ -224,9 +224,11 @@ func (s *Store) Remove(name string) error { } delete(s.items, id) s.vec.Remove(id) - if err := s.writeIndex(); err != nil { - log.Printf("[knowledge] write index error after removing %s: %v", name, err) - } + go func() { + if err := s.writeIndex(); err != nil { + log.Printf("[knowledge] write index error after removing %s: %v", name, err) + } + }() return nil } diff --git a/internal/plugins/cmd/plugin.go b/internal/plugins/cmd/plugin.go index 22d37dc..4e749b0 100644 --- a/internal/plugins/cmd/plugin.go +++ b/internal/plugins/cmd/plugin.go @@ -134,10 +134,6 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error { 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) if timeoutStr == "" { timeoutStr = p.defaultTimeout @@ -218,6 +214,27 @@ func (p *Plugin) Start(s *sdk.PluginSDK) error { }, 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 } @@ -250,9 +267,9 @@ func isDangerousCommand(command string) (bool, string) { pidStr := fmt.Sprintf("%d", pid) lower := strings.ToLower(command) - dangerousPatterns := []string{ - "homed", - pidStr, + selfTargets := []string{"homed", pidStr} + killCmds := []string{"kill ", "killall ", "pkill ", "kill -", "kill -9"} + serviceCmds := []string{ "systemctl stop homeagent", "systemctl restart homeagent", "systemctl kill homeagent", @@ -260,20 +277,19 @@ func isDangerousCommand(command string) (bool, string) { "service homeagent restart", } - killPatterns := []string{"kill ", "killall ", "pkill ", "kill -", "kill -9"} - - for _, kp := range killPatterns { - if strings.Contains(lower, kp) { - for _, target := range dangerousPatterns { - if strings.Contains(lower, strings.ToLower(target)) { - return true, fmt.Sprintf("命令被拦截:禁止向 homed 进程(pid=%s)发送信号", pidStr) - } + for _, kp := range killCmds { + if !strings.Contains(lower, kp) { + continue + } + for _, t := range selfTargets { + if strings.Contains(lower, strings.ToLower(t)) { + return true, fmt.Sprintf("命令被拦截:禁止向 homed 进程(pid=%s)发送信号", pidStr) } } } - for _, dp := range dangerousPatterns[2:] { - if strings.Contains(lower, dp) { + for _, sc := range serviceCmds { + if strings.Contains(lower, sc) { return true, fmt.Sprintf("命令被拦截:禁止操作 homed 服务进程(pid=%s)", pidStr) } }