From 75163f9269a9016c634f7c1d0e700a9214399b87 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=85=AC=E6=98=8E?= <83812544+Ed1s0nZ@users.noreply.github.com> Date: Wed, 22 Jul 2026 10:37:19 +0800 Subject: [PATCH] Add files via upload --- internal/agent/agent.go | 39 +++++---------- internal/agent/agent_test.go | 92 ++++++++++++++++++++++++++++++++++++ 2 files changed, 104 insertions(+), 27 deletions(-) diff --git a/internal/agent/agent.go b/internal/agent/agent.go index e21aaa56..9ad42282 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -24,18 +24,17 @@ import ( // Agent AI代理 type Agent struct { - openAIClient *openai.Client - config *config.OpenAIConfig - agentConfig *config.AgentConfig - mcpServer *mcp.Server - externalMCPMgr *mcp.ExternalMCPManager // 外部MCP管理器 - logger *zap.Logger - maxIterations int - mu sync.RWMutex // 添加互斥锁以支持并发更新 - toolNameMapping map[string]string // 工具名称映射:OpenAI格式 -> 原始格式(用于外部MCP工具) - currentConversationID string // 当前对话ID(用于自动传递给工具) - promptBaseDir string // 解析 system_prompt_path 时相对路径的基准目录(通常为 config.yaml 所在目录) - toolDescriptionMode string // 工具描述模式: "short" | "full",默认 short + openAIClient *openai.Client + config *config.OpenAIConfig + agentConfig *config.AgentConfig + mcpServer *mcp.Server + externalMCPMgr *mcp.ExternalMCPManager // 外部MCP管理器 + logger *zap.Logger + maxIterations int + mu sync.RWMutex // 添加互斥锁以支持并发更新 + toolNameMapping map[string]string // 工具名称映射:OpenAI格式 -> 原始格式(用于外部MCP工具) + promptBaseDir string // 解析 system_prompt_path 时相对路径的基准目录(通常为 config.yaml 所在目录) + toolDescriptionMode string // 工具描述模式: "short" | "full",默认 short } type agentConversationIDKey struct{} @@ -526,12 +525,6 @@ func (a *Agent) executeToolViaMCP(ctx context.Context, toolName string, args map // 如果是record_vulnerability工具,自动添加conversation_id if toolName == builtin.ToolRecordVulnerability { conversationID := agentConversationIDFromContext(ctx) - if conversationID == "" { - a.mu.RLock() - conversationID = a.currentConversationID - a.mu.RUnlock() - } - if conversationID != "" { args["conversation_id"] = conversationID a.logger.Debug("自动添加conversation_id到record_vulnerability工具", @@ -769,16 +762,8 @@ func (a *Agent) ToolsForRole(roleTools []string) []Tool { // ExecuteMCPToolForConversation 在指定会话上下文中执行 MCP 工具(行为与主 Agent 循环中的工具调用一致,如自动注入 conversation_id)。 func (a *Agent) ExecuteMCPToolForConversation(ctx context.Context, conversationID, toolName string, args map[string]interface{}) (*ToolExecutionResult, error) { - a.mu.Lock() - prev := a.currentConversationID - a.currentConversationID = conversationID - a.mu.Unlock() - defer func() { - a.mu.Lock() - a.currentConversationID = prev - a.mu.Unlock() - }() ctx = withAgentConversationID(ctx, conversationID) + ctx = mcp.WithMCPConversationID(ctx, conversationID) return a.executeToolViaMCP(ctx, toolName, args) } diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go index 6e1fa018..34e40512 100644 --- a/internal/agent/agent_test.go +++ b/internal/agent/agent_test.go @@ -3,11 +3,13 @@ package agent import ( "context" "strings" + "sync" "testing" "time" "cyberstrike-ai/internal/config" "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" "go.uber.org/zap" ) @@ -116,3 +118,93 @@ func TestAgentCancelRunningMCPToolsForConversation(t *testing.T) { } t.Fatal("conv-1 execution did not become cancelled") } + +func TestExecuteMCPToolForConversationInjectsConversationID(t *testing.T) { + ag := setupTestAgent(t) + gotArgs := make(chan map[string]interface{}, 1) + ag.mcpServer.RegisterTool(mcp.Tool{Name: builtin.ToolRecordVulnerability, InputSchema: map[string]interface{}{"type": "object"}}, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + gotArgs <- args + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "ok"}}}, nil + }) + + result, err := ag.ExecuteMCPToolForConversation(context.Background(), "conv-record", builtin.ToolRecordVulnerability, map[string]interface{}{}) + if err != nil { + t.Fatalf("ExecuteMCPToolForConversation: %v", err) + } + if result == nil || result.IsError { + t.Fatalf("expected successful result, got %#v", result) + } + + select { + case args := <-gotArgs: + if got := args["conversation_id"]; got != "conv-record" { + t.Fatalf("conversation_id = %#v, want conv-record", got) + } + case <-time.After(time.Second): + t.Fatal("tool was not called") + } +} + +func TestExecuteMCPToolForConversationBindsExecutionConversation(t *testing.T) { + ag := setupTestAgent(t) + ag.mcpServer.ConfigureToolWaitTimeoutSeconds(1) + release := make(chan struct{}) + ag.mcpServer.RegisterTool(mcp.Tool{Name: "slow-bind", InputSchema: map[string]interface{}{"type": "object"}}, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + select { + case <-release: + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "done"}}}, nil + case <-ctx.Done(): + return nil, ctx.Err() + } + }) + + result, err := ag.ExecuteMCPToolForConversation(context.Background(), "conv-bound", "slow-bind", nil) + if err != nil { + t.Fatalf("ExecuteMCPToolForConversation: %v", err) + } + if result == nil || !result.IsError || result.ExecutionID == "" { + t.Fatalf("expected bounded wait result with execution id, result=%#v", result) + } + + exec, ok := ag.mcpServer.GetExecution(result.ExecutionID) + if !ok || exec == nil { + t.Fatalf("missing execution %q", result.ExecutionID) + } + if exec.ConversationID != "conv-bound" { + t.Fatalf("execution conversation = %q, want conv-bound", exec.ConversationID) + } + close(release) +} + +func TestExecuteMCPToolForConversationConcurrentRecordIsolation(t *testing.T) { + ag := setupTestAgent(t) + seen := make(chan string, 2) + ag.mcpServer.RegisterTool(mcp.Tool{Name: builtin.ToolRecordVulnerability, InputSchema: map[string]interface{}{"type": "object"}}, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + if conv, _ := args["conversation_id"].(string); conv != "" { + seen <- conv + } + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "ok"}}}, nil + }) + + var wg sync.WaitGroup + for _, conv := range []string{"conv-a", "conv-b"} { + conv := conv + wg.Add(1) + go func() { + defer wg.Done() + if _, err := ag.ExecuteMCPToolForConversation(context.Background(), conv, builtin.ToolRecordVulnerability, map[string]interface{}{}); err != nil { + t.Errorf("ExecuteMCPToolForConversation %s: %v", conv, err) + } + }() + } + wg.Wait() + close(seen) + + got := map[string]int{} + for conv := range seen { + got[conv]++ + } + if got["conv-a"] != 1 || got["conv-b"] != 1 { + t.Fatalf("conversation ids = %#v, want one call for conv-a and conv-b", got) + } +}