Add files via upload

This commit is contained in:
公明
2026-07-22 10:37:19 +08:00
committed by GitHub
parent 06a9cea97d
commit 75163f9269
2 changed files with 104 additions and 27 deletions
+12 -27
View File
@@ -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)
}
+92
View File
@@ -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)
}
}