From b07a645d34b42c573997b2329ad1767b003ad129 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=85=AC=E6=98=8E?= <83812544+Ed1s0nZ@users.noreply.github.com> Date: Sat, 15 Aug 2026 02:09:57 +0800 Subject: [PATCH] Add files via upload --- internal/app/app.go | 2256 +++++++++++++++++ internal/app/asset_tools.go | 511 ++++ internal/app/asset_tools_test.go | 201 ++ internal/app/c2_hitl_bridge.go | 228 ++ internal/app/c2_lifecycle.go | 104 + internal/app/c2_tools.go | 919 +++++++ internal/app/c2_tools_test.go | 69 + internal/app/cors_security_test.go | 105 + internal/app/main_server_http_redirect.go | 213 ++ .../app/main_server_http_redirect_test.go | 150 ++ internal/app/main_server_tls.go | 86 + internal/app/mcp_authorization.go | 407 +++ internal/app/mcp_authorization_test.go | 261 ++ internal/app/mcp_http_auth_test.go | 59 + internal/app/mcp_project_scope.go | 26 + internal/app/project_fact_tools.go | 389 +++ internal/app/vision_tools.go | 13 + internal/app/vulnerability_tools.go | 466 ++++ internal/einoobserve/attach.go | 455 ++++ internal/einoobserve/attach_test.go | 49 + internal/einoobserve/otel.go | 111 + internal/openai/claude_bridge.go | 1293 ++++++++++ internal/openai/claude_reasoning_roundtrip.go | 81 + .../openai/claude_reasoning_roundtrip_test.go | 211 ++ internal/openai/eino_sse_sanitizer.go | 149 ++ internal/openai/eino_sse_sanitizer_test.go | 303 +++ .../openai/normalize_streaming_delta_test.go | 56 + internal/openai/openai.go | 616 +++++ internal/openai/reasoning_payload.go | 117 + internal/openai/reasoning_payload_test.go | 250 ++ .../openai/reasoning_tool_choice_compat.go | 69 + internal/openai/sse_stream.go | 20 + internal/openai/summarization_diag.go | 108 + internal/openai/summarization_diag_test.go | 47 + internal/tooloutput/spill.go | 292 +++ internal/tooloutput/spill_test.go | 58 + 36 files changed, 10748 insertions(+) create mode 100644 internal/app/app.go create mode 100644 internal/app/asset_tools.go create mode 100644 internal/app/asset_tools_test.go create mode 100644 internal/app/c2_hitl_bridge.go create mode 100644 internal/app/c2_lifecycle.go create mode 100644 internal/app/c2_tools.go create mode 100644 internal/app/c2_tools_test.go create mode 100644 internal/app/cors_security_test.go create mode 100644 internal/app/main_server_http_redirect.go create mode 100644 internal/app/main_server_http_redirect_test.go create mode 100644 internal/app/main_server_tls.go create mode 100644 internal/app/mcp_authorization.go create mode 100644 internal/app/mcp_authorization_test.go create mode 100644 internal/app/mcp_http_auth_test.go create mode 100644 internal/app/mcp_project_scope.go create mode 100644 internal/app/project_fact_tools.go create mode 100644 internal/app/vision_tools.go create mode 100644 internal/app/vulnerability_tools.go create mode 100644 internal/einoobserve/attach.go create mode 100644 internal/einoobserve/attach_test.go create mode 100644 internal/einoobserve/otel.go create mode 100644 internal/openai/claude_bridge.go create mode 100644 internal/openai/claude_reasoning_roundtrip.go create mode 100644 internal/openai/claude_reasoning_roundtrip_test.go create mode 100644 internal/openai/eino_sse_sanitizer.go create mode 100644 internal/openai/eino_sse_sanitizer_test.go create mode 100644 internal/openai/normalize_streaming_delta_test.go create mode 100644 internal/openai/openai.go create mode 100644 internal/openai/reasoning_payload.go create mode 100644 internal/openai/reasoning_payload_test.go create mode 100644 internal/openai/reasoning_tool_choice_compat.go create mode 100644 internal/openai/sse_stream.go create mode 100644 internal/openai/summarization_diag.go create mode 100644 internal/openai/summarization_diag_test.go create mode 100644 internal/tooloutput/spill.go create mode 100644 internal/tooloutput/spill_test.go diff --git a/internal/app/app.go b/internal/app/app.go new file mode 100644 index 00000000..7cf5b079 --- /dev/null +++ b/internal/app/app.go @@ -0,0 +1,2256 @@ +package app + +import ( + "context" + "crypto/subtle" + "crypto/tls" + "database/sql" + "fmt" + "net" + "net/http" + "net/url" + "os" + "path/filepath" + "strings" + "sync" + "time" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/audit" + "cyberstrike-ai/internal/authctx" + "cyberstrike-ai/internal/c2" + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/einoobserve" + "cyberstrike-ai/internal/handler" + "cyberstrike-ai/internal/hitl" + "cyberstrike-ai/internal/knowledge" + "cyberstrike-ai/internal/logger" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" + "cyberstrike-ai/internal/monitor" + "cyberstrike-ai/internal/multiagent" + "cyberstrike-ai/internal/robot" + "cyberstrike-ai/internal/security" + "cyberstrike-ai/internal/skillpackage" + + "github.com/gin-gonic/gin" + "github.com/google/uuid" + "go.uber.org/zap" + "golang.org/x/net/http2" +) + +// App 应用 +type App struct { + config *config.Config + logger *logger.Logger + router *gin.Engine + mcpServer *mcp.Server + externalMCPMgr *mcp.ExternalMCPManager + agent *agent.Agent + executor *security.Executor + db *database.DB + knowledgeDB *database.DB // 知识库数据库连接(如果使用独立数据库) + auth *security.AuthManager + knowledgeManager *knowledge.Manager // 知识库管理器(用于动态初始化) + knowledgeRetriever *knowledge.Retriever // 知识库检索器(用于动态初始化) + knowledgeIndexer *knowledge.Indexer // 知识库索引器(用于动态初始化) + knowledgeHandler *handler.KnowledgeHandler // 知识库处理器(用于动态初始化) + agentHandler *handler.AgentHandler // Agent处理器(用于更新知识库管理器) + robotHandler *handler.RobotHandler // 机器人处理器(钉钉/飞书/企业微信等) + robotMu sync.Mutex // 保护机器人长连接的 cancel + dingCancel context.CancelFunc // 钉钉 Stream 取消函数,用于配置变更时重启 + larkCancel context.CancelFunc // 飞书长连接取消函数,用于配置变更时重启 + wechatCancel context.CancelFunc // 微信 iLink 长轮询取消函数 + telegramCancel context.CancelFunc // Telegram 长轮询取消函数 + slackCancel context.CancelFunc // Slack Socket Mode 取消函数 + discordCancel context.CancelFunc // Discord Gateway 取消函数 + qqCancel context.CancelFunc // QQ WebSocket 取消函数 + alertCancel context.CancelFunc // 漏洞提醒持久化投递 worker + c2Manager *c2.Manager // C2 管理器(未启用 C2 时为 nil) + c2Watchdog *c2.SessionWatchdog // C2 会话看门狗 + c2WatchdogCancel context.CancelFunc // 看门狗取消函数 + c2Handler *handler.C2Handler // C2 REST(与 Manager 生命周期同步) + auditSvc *audit.Service +} + +// New 创建新应用 +func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error) { + if err := multiagent.InitADK(); err != nil { + return nil, fmt.Errorf("初始化 Eino ADK: %w", err) + } + + gin.SetMode(gin.ReleaseMode) + router := gin.Default() + + // CORS中间件 + router.Use(corsMiddleware(cfg.Server.CORSAllowedOrigins)) + + // 初始化数据库 + dbPath := cfg.Database.Path + if dbPath == "" { + dbPath = "data/conversations.db" + } + + // 确保目录存在 + if err := os.MkdirAll(filepath.Dir(dbPath), 0755); err != nil { + return nil, fmt.Errorf("创建数据库目录失败: %w", err) + } + + db, err := database.NewDB(dbPath, log.Logger) + if err != nil { + return nil, fmt.Errorf("初始化数据库失败: %w", err) + } + + // 认证管理器(数据库初始化后挂载 RBAC) + authManager := security.NewAuthManager(cfg.Auth.SessionDurationHours) + if generatedPassword, err := authManager.AttachRBACStore(db); err != nil { + return nil, fmt.Errorf("初始化RBAC失败: %w", err) + } else if generatedPassword != "" { + config.PrintBootstrapAdminPassword(generatedPassword) + } + for platform, userID := range cfg.Robots.ServiceAccountUserIDs() { + user, userErr := db.GetRBACUserByID(userID) + if userErr != nil || !user.Enabled { + return nil, fmt.Errorf("robots.%s.auth.service_user_id 必须指向已启用的 RBAC 用户", platform) + } + } + + auditSvc := audit.NewService(db, cfg, log.Logger) + audit.RegisterConversationCreateHook(auditSvc) + auditSvc.PurgeExpired() + audit.StartRetentionLoop(auditSvc, log.Logger) + if err := db.PurgeWorkflowPackageLifecycle(time.Now().UTC()); err != nil { + log.Logger.Warn("清理过期工作流包记录失败", zap.Error(err)) + } + go func() { + ticker := time.NewTicker(time.Hour) + defer ticker.Stop() + for range ticker.C { + if err := db.PurgeWorkflowPackageLifecycle(time.Now().UTC()); err != nil { + log.Logger.Warn("清理过期工作流包记录失败", zap.Error(err)) + } + } + }() + + monitorRetention := monitor.NewService(db, cfg, log.Logger) + monitorRetention.PurgeExpired() + monitor.StartRetentionLoop(monitorRetention, log.Logger) + + if err := handler.NewHITLManager(db, log.Logger).EnsureSchema(); err != nil { + log.Logger.Warn("初始化 HITL 表失败", zap.Error(err)) + } + hitlRetention := hitl.NewService(db, cfg, log.Logger) + hitlRetention.PurgeExpired() + hitl.StartRetentionLoop(hitlRetention, log.Logger) + + // 创建MCP服务器(带数据库持久化) + mcpServer := mcp.NewServerWithStorage(log.Logger, db) + mcpServer.SetToolAuthorizer(mcpToolAuthorizer(db)) + mcpServer.ConfigureHTTPToolCallTimeoutFromAgentMinutes(cfg.Agent.ToolTimeoutMinutes) + mcpServer.ConfigureToolWaitTimeoutSeconds(cfg.Agent.ToolWaitTimeoutSeconds) + mcpServer.ConfigureToolResultMaxBytes(cfg.MultiAgent.EinoMiddleware.ReductionMaxLengthForTruncEffective()) + mcpServer.ConfigureToolResultSpillRoot(cfg.MultiAgent.EinoMiddleware.ReductionRootDir) + + // 创建安全工具执行器 + executor := security.NewExecutor(&cfg.Security, mcpServer, log.Logger) + executor.SetShellNoOutputTimeoutSeconds(cfg.Agent.ShellNoOutputTimeoutSeconds) + executor.SetToolOutputMaxBytes(cfg.MultiAgent.EinoMiddleware.ReductionMaxLengthForTruncEffective()) + executor.SetToolOutputSpillRoot(cfg.MultiAgent.EinoMiddleware.ReductionRootDir) + + // 注册工具 + executor.RegisterTools(mcpServer) + + // 注册漏洞记录工具 + registerVulnerabilityTools(mcpServer, db, log.Logger) + registerAssetTools(mcpServer, db, log.Logger) + registerProjectFactTools(mcpServer, db, cfg, log.Logger) + registerVisionTools(mcpServer, cfg, log.Logger) + + // 创建外部MCP管理器(使用与内部MCP服务器相同的存储) + externalMCPMgr := mcp.NewExternalMCPManagerWithStorage(log.Logger, db) + externalMCPMgr.SetToolAuthorizer(externalMCPToolAuthorizer()) + externalMCPMgr.ConfigureToolWaitTimeoutSeconds(cfg.Agent.ToolWaitTimeoutSeconds) + externalMCPMgr.ConfigureToolResultMaxBytes(cfg.MultiAgent.EinoMiddleware.ReductionMaxLengthForTruncEffective()) + externalMCPMgr.ConfigureToolResultSpillRoot(cfg.MultiAgent.EinoMiddleware.ReductionRootDir) + externalMCPMgr.ConfigureResilience(mcp.ExternalMCPResilienceConfig{ + MaxConcurrentPerServer: cfg.Agent.ExternalMCPMaxConcurrentPerServer, + MaxConcurrentTotal: cfg.Agent.ExternalMCPMaxConcurrentTotal, + CircuitFailureThreshold: cfg.Agent.ExternalMCPCircuitFailureThreshold, + CircuitCooldown: time.Duration(cfg.Agent.ExternalMCPCircuitCooldownSeconds) * time.Second, + }) + mcp.RegisterExecutionControlTools(mcpServer, externalMCPMgr) + if cfg.ExternalMCP.Servers != nil { + externalMCPMgr.LoadConfigs(&cfg.ExternalMCP) + // 启动所有启用的外部MCP客户端 + externalMCPMgr.StartAllEnabled() + } + + execReconciler := monitor.NewExecutionReconciler(db, mcpServer, externalMCPMgr, log.Logger) + execReconciler.ReconcileOnStartup() + monitor.StartStaleRunningReconcileLoop(execReconciler, log.Logger) + + // 创建Agent + maxIterations := cfg.Agent.MaxIterations + if maxIterations <= 0 { + maxIterations = 30 // 默认值 + } + agent := agent.NewAgent(&cfg.OpenAI, &cfg.Agent, mcpServer, externalMCPMgr, log.Logger, maxIterations) + agent.UpdateToolDescriptionMode(cfg.Security.ToolDescriptionMode) + + // 初始化知识库模块(如果启用) + var knowledgeManager *knowledge.Manager + var knowledgeRetriever *knowledge.Retriever + var knowledgeIndexer *knowledge.Indexer + var knowledgeHandler *handler.KnowledgeHandler + + var knowledgeDBConn *database.DB + log.Logger.Debug("检查知识库配置", zap.Bool("enabled", cfg.Knowledge.Enabled)) + if cfg.Knowledge.Enabled { + // 确定知识库数据库路径 + knowledgeDBPath := cfg.Database.KnowledgeDBPath + var knowledgeDB *sql.DB + + if knowledgeDBPath != "" { + // 使用独立的知识库数据库 + // 确保目录存在 + if err := os.MkdirAll(filepath.Dir(knowledgeDBPath), 0755); err != nil { + return nil, fmt.Errorf("创建知识库数据库目录失败: %w", err) + } + + var err error + knowledgeDBConn, err = database.NewKnowledgeDB(knowledgeDBPath, log.Logger) + if err != nil { + return nil, fmt.Errorf("初始化知识库数据库失败: %w", err) + } + knowledgeDB = knowledgeDBConn.DB + log.Logger.Info("使用独立的知识库数据库", zap.String("path", knowledgeDBPath)) + } else { + // 向后兼容:使用会话数据库 + knowledgeDB = db.DB + log.Logger.Info("使用会话数据库存储知识库数据(建议配置knowledge_db_path以分离数据)") + } + + // 创建知识库管理器 + knowledgeManager = knowledge.NewManager(knowledgeDB, cfg.Knowledge.BasePath, log.Logger) + + // 创建嵌入器 + // 使用OpenAI配置的API Key(如果知识库配置中没有指定) + if cfg.Knowledge.Embedding.APIKey == "" { + cfg.Knowledge.Embedding.APIKey = cfg.OpenAI.APIKey + } + if cfg.Knowledge.Embedding.BaseURL == "" { + cfg.Knowledge.Embedding.BaseURL = cfg.OpenAI.BaseURL + } + + embedder, err := knowledge.NewEmbedder(context.Background(), &cfg.Knowledge, &cfg.OpenAI, log.Logger) + if err != nil { + return nil, fmt.Errorf("初始化知识库嵌入器失败: %w", err) + } + + // 创建检索器(Eino MultiQuery + 重排流水线) + retrievalConfig := knowledge.RetrievalConfigFromYAML(cfg.Knowledge.Retrieval) + knowledgeRetriever = knowledge.NewRetriever(knowledgeDB, embedder, retrievalConfig, log.Logger) + if err := knowledge.WireRetrieverPipeline(context.Background(), knowledgeRetriever, &cfg.OpenAI); err != nil { + return nil, fmt.Errorf("初始化知识库检索流水线失败: %w", err) + } + + // 创建索引器(Eino Compose 链) + knowledgeIndexer, err = knowledge.NewIndexer(context.Background(), knowledgeDB, embedder, log.Logger, &cfg.Knowledge) + if err != nil { + return nil, fmt.Errorf("初始化知识库索引器失败: %w", err) + } + + // 注册知识检索工具到MCP服务器 + knowledge.RegisterKnowledgeTool(mcpServer, knowledgeRetriever, knowledgeManager, log.Logger) + + // 创建知识库API处理器 + knowledgeHandler = handler.NewKnowledgeHandler(knowledgeManager, knowledgeRetriever, knowledgeIndexer, db, log.Logger) + knowledgeHandler.SetAudit(auditSvc) + log.Logger.Info("知识库模块初始化完成", zap.Bool("handler_created", knowledgeHandler != nil)) + + // 扫描知识库并建立索引(异步) + go func() { + itemsToIndex, err := knowledgeManager.ScanKnowledgeBase() + if err != nil { + log.Logger.Warn("扫描知识库失败", zap.Error(err)) + return + } + + // 检查是否已有索引 + hasIndex, err := knowledgeIndexer.HasIndex() + if err != nil { + log.Logger.Warn("检查索引状态失败", zap.Error(err)) + return + } + + if hasIndex { + // 如果已有索引,只索引新添加或更新的项 + if len(itemsToIndex) > 0 { + log.Logger.Info("检测到已有知识库索引,开始增量索引", zap.Int("count", len(itemsToIndex))) + ctx := context.Background() + consecutiveFailures := 0 + var firstFailureItemID string + var firstFailureError error + failedCount := 0 + + for _, itemID := range itemsToIndex { + if err := knowledgeIndexer.IndexItem(ctx, itemID); err != nil { + failedCount++ + consecutiveFailures++ + + if consecutiveFailures == 1 { + firstFailureItemID = itemID + firstFailureError = err + log.Logger.Warn("索引知识项失败", zap.String("itemId", itemID), zap.Error(err)) + } + + // 如果连续失败2次,立即停止增量索引 + if consecutiveFailures >= 2 { + log.Logger.Error("连续索引失败次数过多,立即停止增量索引", + zap.Int("consecutiveFailures", consecutiveFailures), + zap.Int("totalItems", len(itemsToIndex)), + zap.String("firstFailureItemId", firstFailureItemID), + zap.Error(firstFailureError), + ) + break + } + continue + } + + // 成功时重置连续失败计数 + if consecutiveFailures > 0 { + consecutiveFailures = 0 + firstFailureItemID = "" + firstFailureError = nil + } + } + log.Logger.Info("增量索引完成", zap.Int("totalItems", len(itemsToIndex)), zap.Int("failedCount", failedCount)) + } else { + log.Logger.Info("检测到已有知识库索引,没有需要索引的新项或更新项") + } + return + } + + // 冷启动:仅为尚无向量的知识项构建索引(与 IndexMissing 语义一致) + log.Logger.Info("未检测到知识库索引,开始自动构建索引") + ctx := context.Background() + if err := knowledgeIndexer.IndexMissing(ctx); err != nil { + log.Logger.Warn("自动构建知识库索引失败", zap.Error(err)) + } + }() + } + + // 配置文件路径必须由入口传入(与 flag -config 一致)。勿再用 os.Args[1],否则 ./cyberstrike-ai --https 会把 --https 当成路径。 + configPath = strings.TrimSpace(configPath) + if configPath == "" { + configPath = "config.yaml" + } + + skillsDir := skillpackage.SkillsRootFromConfig(cfg.SkillsDir, configPath) + log.Logger.Debug("Skills 目录(Eino ADK skill 中间件 + Web 管理 API)", zap.String("skillsDir", skillsDir)) + configDir := filepath.Dir(configPath) + plantaskRel := strings.TrimSpace(cfg.MultiAgent.EinoMiddleware.PlantaskRelDir) + if plantaskRel == "" { + plantaskRel = ".eino/plantask" + } + plantaskBase := filepath.Join(skillsDir, plantaskRel) + // Match eino_adk_run_loop: checkpoint_dir is used as configured (relative to process CWD when not absolute). + checkpointBase := strings.TrimSpace(cfg.MultiAgent.EinoMiddleware.CheckpointDir) + reductionRoot := strings.TrimSpace(cfg.MultiAgent.EinoMiddleware.ReductionRootDir) + workspaceRoot := strings.TrimSpace(cfg.Agent.WorkspaceRootDir) + db.SetEinoConversationDirs(plantaskBase, checkpointBase, reductionRoot, workspaceRoot) + agent.SetPromptBaseDir(configDir) + + agentsDir := cfg.AgentsDir + if agentsDir == "" { + agentsDir = "agents" + } + if !filepath.IsAbs(agentsDir) { + agentsDir = filepath.Join(configDir, agentsDir) + } + if err := os.MkdirAll(agentsDir, 0755); err != nil { + log.Logger.Warn("创建 agents 目录失败", zap.String("path", agentsDir), zap.Error(err)) + } + markdownAgentsHandler := handler.NewMarkdownAgentsHandler(agentsDir) + markdownAgentsHandler.SetAudit(auditSvc) + log.Logger.Debug("多代理 Markdown 子 Agent 目录", zap.String("agentsDir", agentsDir)) + + // 创建处理器 + agentHandler := handler.NewAgentHandler(agent, db, cfg, log.Logger) + agentHandler.SetAudit(auditSvc) + agentHandler.SetAgentsMarkdownDir(agentsDir) + // 如果知识库已启用,设置知识库管理器到AgentHandler以便记录检索日志 + if knowledgeManager != nil { + agentHandler.SetKnowledgeManager(knowledgeManager) + } + monitorHandler := handler.NewMonitorHandler(mcpServer, executor, db, log.Logger) + monitorHandler.SetAudit(auditSvc) + monitorHandler.SetMonitorRetention(monitorRetention) + monitorHandler.SetExternalMCPManager(externalMCPMgr) // 设置外部MCP管理器,以便获取外部MCP执行记录 + monitorHandler.SetTaskManager(agentHandler.TaskManager()) + monitorHandler.SetAgentHandler(agentHandler) + notificationHandler := handler.NewNotificationHandler(db, agentHandler, log.Logger) + groupHandler := handler.NewGroupHandler(db, log.Logger) + authHandler := handler.NewAuthHandler(authManager, cfg, configPath, log.Logger) + authHandler.SetAudit(auditSvc) + attackChainHandler := handler.NewAttackChainHandler(db, &cfg.OpenAI, log.Logger) + vulnerabilityHandler := handler.NewVulnerabilityHandler(db, log.Logger) + assetHandler := handler.NewAssetHandler(db, log.Logger) + projectHandler := handler.NewProjectHandler(db, log.Logger) + rbacHandler := handler.NewRBACHandler(db, log.Logger) + rbacHandler.SetAudit(auditSvc) + rbacHandler.SetAuthManager(authManager) + workflowHandler := handler.NewWorkflowHandler(db, log.Logger) + workflowHandler.SetAudit(auditSvc) + workflowHandler.SetRuntime(agent, cfg) + vulnerabilityHandler.SetAudit(auditSvc) + webshellHandler := handler.NewWebShellHandler(log.Logger, db) + webshellHandler.SetAudit(auditSvc) + chatUploadsHandler := handler.NewChatUploadsHandler(log.Logger, db) + chatUploadsHandler.SetAudit(auditSvc) + registerWebshellTools(mcpServer, db, webshellHandler, log.Logger) + registerWebshellManagementTools(mcpServer, db, webshellHandler, log.Logger) + configHandler := handler.NewConfigHandler(configPath, cfg, mcpServer, executor, agent, attackChainHandler, externalMCPMgr, log.Logger) + configHandler.SetDB(db) + configHandler.SetAudit(auditSvc) + agentHandler.SetHitlToolWhitelistSaver(configHandler) + agentHandler.SetHitlAuditStrategySaver(configHandler) + agentHandler.SetHitlDefaultReviewerSaver(configHandler) + externalMCPHandler := handler.NewExternalMCPHandler(externalMCPMgr, cfg, configPath, log.Logger) + externalMCPHandler.SetAudit(auditSvc) + roleHandler := handler.NewRoleHandler(cfg, configPath, log.Logger) + roleHandler.SetAudit(auditSvc) + skillsHandler := handler.NewSkillsHandler(cfg, configPath, log.Logger) + skillsHandler.SetAudit(auditSvc) + fofaHandler := handler.NewFofaHandler(cfg, log.Logger) + terminalHandler := handler.NewTerminalHandler(log.Logger) + if db != nil { + skillsHandler.SetDB(db) // 设置数据库连接以便获取调用统计 + } + + // ============================================================================ + // 初始化 C2 模块(可按配置关闭,节省本机部署资源) + // ============================================================================ + c2Manager, c2Watchdog, watchdogCancel := setupC2Runtime(cfg, db, agentHandler, log.Logger) + if c2Manager != nil { + registerC2Tools(mcpServer, c2Manager, log.Logger, cfg.Server.Port) + } + c2Handler := handler.NewC2Handler(c2Manager, log.Logger) + c2Handler.SetAudit(auditSvc) + + // 创建OpenAPI处理器 + conversationHandler := handler.NewConversationHandler(db, log.Logger) + conversationHandler.SetAudit(auditSvc) + conversationHandler.SetTaskStopper(agentHandler) + conversationHandler.SetTaskStateProvider(agentHandler) + auditHandler := handler.NewAuditHandler(db, auditSvc, log.Logger) + robotHandler := handler.NewRobotHandler(cfg, db, agentHandler, log.Logger) + robotHandler.SetAudit(auditSvc) + db.SetVulnerabilityCreatedHook(robotHandler.NotifyNewVulnerability) + openAPIHandler := handler.NewOpenAPIHandler(db, log.Logger, conversationHandler, agentHandler) + + // 创建 App 实例(部分字段稍后填充) + app := &App{ + config: cfg, + logger: log, + router: router, + mcpServer: mcpServer, + externalMCPMgr: externalMCPMgr, + agent: agent, + executor: executor, + db: db, + knowledgeDB: knowledgeDBConn, + auth: authManager, + knowledgeManager: knowledgeManager, + knowledgeRetriever: knowledgeRetriever, + knowledgeIndexer: knowledgeIndexer, + knowledgeHandler: knowledgeHandler, + agentHandler: agentHandler, + robotHandler: robotHandler, + c2Manager: c2Manager, + c2Watchdog: c2Watchdog, + c2WatchdogCancel: watchdogCancel, + c2Handler: c2Handler, + auditSvc: auditSvc, + } + // 飞书/钉钉长连接(无需公网),启用时在后台启动;后续前端应用配置时会通过 RestartRobotConnections 重启 + app.startRobotConnections() + alertCtx, alertCancel := context.WithCancel(context.Background()) + app.alertCancel = alertCancel + go robotHandler.RunVulnerabilityAlertWorker(alertCtx) + + // 设置漏洞工具注册器(内置工具,必须设置) + vulnerabilityRegistrar := func() error { + registerVulnerabilityTools(mcpServer, db, log.Logger) + registerAssetTools(mcpServer, db, log.Logger) + registerProjectFactTools(mcpServer, db, cfg, log.Logger) + registerVisionTools(mcpServer, cfg, log.Logger) + return nil + } + configHandler.SetVulnerabilityToolRegistrar(vulnerabilityRegistrar) + + // 设置 WebShell 工具注册器(ApplyConfig 时重新注册) + webshellRegistrar := func() error { + registerWebshellTools(mcpServer, db, webshellHandler, log.Logger) + registerWebshellManagementTools(mcpServer, db, webshellHandler, log.Logger) + return nil + } + configHandler.SetWebshellToolRegistrar(webshellRegistrar) + + // Skills 由 Eino ADK skill 中间件提供(多代理);此处不注册 MCP 形态的技能工具 + configHandler.SetSkillsToolRegistrar(func() error { return nil }) + + handler.RegisterBatchTaskMCPTools(mcpServer, agentHandler, log.Logger) + batchTaskToolRegistrar := func() error { + handler.RegisterBatchTaskMCPTools(mcpServer, agentHandler, log.Logger) + return nil + } + configHandler.SetBatchTaskToolRegistrar(batchTaskToolRegistrar) + + // 设置知识库初始化器(用于动态初始化,需要在 App 创建后设置) + configHandler.SetKnowledgeInitializer(func() (*handler.KnowledgeHandler, error) { + knowledgeHandler, err := initializeKnowledge(cfg, db, knowledgeDBConn, mcpServer, agentHandler, app, log.Logger) + if err != nil { + return nil, err + } + + // 动态初始化后,设置知识库工具注册器和检索器更新器 + // 这样后续 ApplyConfig 时就能重新注册工具了 + if app.knowledgeRetriever != nil && app.knowledgeManager != nil { + // 创建闭包,捕获knowledgeRetriever和knowledgeManager的引用 + registrar := func() error { + knowledge.RegisterKnowledgeTool(mcpServer, app.knowledgeRetriever, app.knowledgeManager, log.Logger) + return nil + } + configHandler.SetKnowledgeToolRegistrar(registrar) + // 设置检索器更新器,以便在ApplyConfig时更新检索器配置 + configHandler.SetRetrieverUpdater(app.knowledgeRetriever) + log.Logger.Info("动态初始化后已设置知识库工具注册器和检索器更新器") + } + + return knowledgeHandler, nil + }) + + // 如果知识库已启用,设置知识库工具注册器和检索器更新器 + if cfg.Knowledge.Enabled && knowledgeRetriever != nil && knowledgeManager != nil { + // 创建闭包,捕获knowledgeRetriever和knowledgeManager的引用 + registrar := func() error { + knowledge.RegisterKnowledgeTool(mcpServer, knowledgeRetriever, knowledgeManager, log.Logger) + return nil + } + configHandler.SetKnowledgeToolRegistrar(registrar) + // 设置检索器更新器,以便在ApplyConfig时更新检索器配置 + configHandler.SetRetrieverUpdater(knowledgeRetriever) + } + + // 设置机器人连接重启器,前端应用配置后无需重启服务即可使钉钉/飞书/微信新配置生效 + configHandler.SetRobotRestarter(app) + + wechatRobotHandler := handler.NewWechatRobotHandler(cfg, configHandler, log.Logger) + + configHandler.SetC2Runtime(app) + configHandler.SetC2ToolRegistrar(func() error { + if app.config.C2.EnabledEffective() && app.c2Manager != nil { + registerC2Tools(mcpServer, app.c2Manager, log.Logger, app.config.Server.Port) + } + return nil + }) + + // 设置路由(使用 App 实例以便动态获取 handler) + setupRoutes( + router, + authHandler, + agentHandler, + monitorHandler, + notificationHandler, + conversationHandler, + robotHandler, + wechatRobotHandler, + groupHandler, + configHandler, + externalMCPHandler, + attackChainHandler, + app, // 传递 App 实例以便动态获取 knowledgeHandler + vulnerabilityHandler, + assetHandler, + projectHandler, + workflowHandler, + webshellHandler, + chatUploadsHandler, + roleHandler, + skillsHandler, + markdownAgentsHandler, + fofaHandler, + terminalHandler, + app.c2Handler, + auditHandler, + auditSvc, + rbacHandler, + mcpServer, + authManager, + openAPIHandler, + ) + + return app, nil + +} + +// mcpHandlerWithAuth 在鉴权通过后转发到 MCP 处理;若配置了 auth_header 则校验请求头,否则直接放行 +func (a *App) mcpHandlerWithAuth(w http.ResponseWriter, r *http.Request) { + cfg := a.config.MCP + if authHeader := strings.TrimSpace(r.Header.Get("Authorization")); len(authHeader) > 7 && strings.EqualFold(authHeader[:7], "Bearer ") { + if session, ok := a.auth.ValidateToken(strings.TrimSpace(authHeader[7:])); ok && session.Permissions["mcp:execute"] { + principal := authctx.NewPrincipalWithScopes(session.UserID, session.Username, session.Scope, session.Permissions, session.PermissionScopes) + a.mcpServer.HandleHTTP(w, r.WithContext(authctx.WithPrincipal(r.Context(), principal))) + return + } + } + if !cfg.AllowGlobalAccess || strings.TrimSpace(cfg.AuthHeader) == "" || strings.TrimSpace(cfg.AuthHeaderValue) == "" { + http.Error(w, "use an authorized user bearer token; global MCP service access is disabled", http.StatusUnauthorized) + return + } + if subtle.ConstantTimeCompare([]byte(r.Header.Get(cfg.AuthHeader)), []byte(cfg.AuthHeaderValue)) != 1 { + a.logger.Logger.Debug("MCP 鉴权失败:header 缺失或值不匹配", zap.String("header", cfg.AuthHeader)) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusUnauthorized) + w.Write([]byte(`{"error":"unauthorized"}`)) + return + } + permissions := make(map[string]bool, len(security.PermissionCatalog)) + for permission := range security.PermissionCatalog { + permissions[permission] = true + } + principal := authctx.NewPrincipal("service:mcp", "mcp-service", database.RBACScopeAll, permissions) + r = r.WithContext(authctx.WithPrincipal(r.Context(), principal)) + a.mcpServer.HandleHTTP(w, r) +} + +// Run 启动应用(向后兼容,不支持优雅关闭) +func (a *App) Run() error { + return a.RunWithContext(context.Background()) +} + +// RunWithContext 启动应用,支持通过 context 取消来优雅关闭 +func (a *App) RunWithContext(ctx context.Context) error { + // 启动MCP服务器(如果启用) + var mcpServer *http.Server + if a.config.MCP.Enabled { + mcpAddr := fmt.Sprintf("%s:%d", a.config.MCP.Host, a.config.MCP.Port) + a.logger.Info("启动MCP服务器", zap.String("address", mcpAddr)) + + mux := http.NewServeMux() + mux.HandleFunc("/mcp", a.mcpHandlerWithAuth) + + mcpServer = &http.Server{Addr: mcpAddr, Handler: mux} + go func() { + if err := mcpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed { + a.logger.Error("MCP服务器启动失败", zap.Error(err)) + } + }() + } + + // 启动主服务器(可选 HTTPS + HTTP/2,见 config server.tls_*) + addr := fmt.Sprintf("%s:%d", a.config.Server.Host, a.config.Server.Port) + tlsMode, tlsConf, certFile, keyFile, tlsErr := prepareMainServerTLS(&a.config.Server) + if tlsErr != nil { + return tlsErr + } + + srv := &http.Server{Addr: addr, Handler: a.router} + var mainMux *mainServerMux + httpRedirect := config.ServerHTTPRedirectEnabled(&a.config.Server) + if tlsMode != mainTLSOff { + srv.TLSConfig = tlsConf + if err := http2.ConfigureServer(srv, &http2.Server{}); err != nil { + return fmt.Errorf("主服务 HTTP/2 配置失败: %w", err) + } + switch tlsMode { + case mainTLSFromFiles: + a.logger.Debug("启动 HTTPS 主服务(已启用 HTTP/2 协商)", + zap.String("address", addr), + zap.String("cert", certFile), + ) + case mainTLSInMemorySelfSigned: + a.logger.Debug("启动 HTTPS 主服务(内存自签证书,仅测试;已启用 HTTP/2 协商)", + zap.String("address", addr), + ) + } + if httpRedirect { + a.logger.Debug("已启用 HTTP→HTTPS 自动跳转(同端口嗅探分流)", zap.String("address", addr)) + } + } else { + a.logger.Debug("启动 HTTP 主服务", zap.String("address", addr)) + } + + // 监听 context 取消,优雅关闭 HTTP 服务器 + go func() { + <-ctx.Done() + shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + if mainMux != nil { + if err := mainMux.Shutdown(shutdownCtx); err != nil { + a.logger.Error("HTTP/HTTPS 分流服务器关闭失败", zap.Error(err)) + } + } else if err := srv.Shutdown(shutdownCtx); err != nil { + a.logger.Error("HTTP服务器关闭失败", zap.Error(err)) + } + if mcpServer != nil { + if err := mcpServer.Shutdown(shutdownCtx); err != nil { + a.logger.Error("MCP服务器关闭失败", zap.Error(err)) + } + } + }() + + var err error + switch { + case tlsMode != mainTLSOff && httpRedirect: + var tlsConfReady *tls.Config + tlsConfReady, err = ensureMainTLSConfigCerts(tlsMode, tlsConf, certFile, keyFile) + if err != nil { + return fmt.Errorf("加载 TLS 证书: %w", err) + } + srv.TLSConfig = tlsConfReady + var ln net.Listener + ln, err = net.Listen("tcp", addr) + if err != nil { + return err + } + mainMux = newMainServerMux(ln, srv, portFromListenAddr(addr), a.logger.Logger) + err = mainMux.Serve() + case tlsMode == mainTLSOff: + err = srv.ListenAndServe() + case tlsMode == mainTLSFromFiles: + err = srv.ListenAndServeTLS(certFile, keyFile) + case tlsMode == mainTLSInMemorySelfSigned: + var ln net.Listener + ln, err = tls.Listen("tcp", addr, srv.TLSConfig) + if err == nil { + err = srv.Serve(ln) + } + default: + err = srv.ListenAndServe() + } + if err != nil && err != http.ErrServerClosed { + return err + } + return nil +} + +// Shutdown 关闭应用 +func (a *App) Shutdown() { + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 5*time.Second) + _ = einoobserve.ShutdownOtel(shutdownCtx) + shutdownCancel() + if a.alertCancel != nil { + a.alertCancel() + a.alertCancel = nil + } + + // 停止钉钉/飞书长连接 + a.robotMu.Lock() + if a.dingCancel != nil { + a.dingCancel() + a.dingCancel = nil + } + if a.larkCancel != nil { + a.larkCancel() + a.larkCancel = nil + } + a.robotMu.Unlock() + + a.shutdownC2() + + // 停止所有外部MCP客户端 + if a.externalMCPMgr != nil { + a.externalMCPMgr.StopAll() + } + + // 关闭知识库数据库连接(如果使用独立数据库) + if a.knowledgeDB != nil { + if err := a.knowledgeDB.Close(); err != nil { + a.logger.Logger.Warn("关闭知识库数据库连接失败", zap.Error(err)) + } + } + + // 关闭主数据库连接 + if a.db != nil { + if err := a.db.Close(); err != nil { + a.logger.Logger.Warn("关闭主数据库连接失败", zap.Error(err)) + } + } +} + +// startRobotConnections 根据当前配置启动钉钉/飞书长连接(不先关闭已有连接,仅用于首次启动) +func (a *App) startRobotConnections() { + a.robotMu.Lock() + defer a.robotMu.Unlock() + cfg := a.config + if cfg.Robots.Lark.Enabled && cfg.Robots.Lark.AppID != "" && cfg.Robots.Lark.AppSecret != "" { + ctx, cancel := context.WithCancel(context.Background()) + a.larkCancel = cancel + go robot.StartLark(ctx, cfg.Robots, a.robotHandler, a.logger.Logger) + } + if cfg.Robots.Dingtalk.Enabled && cfg.Robots.Dingtalk.ClientID != "" && cfg.Robots.Dingtalk.ClientSecret != "" { + ctx, cancel := context.WithCancel(context.Background()) + a.dingCancel = cancel + go robot.StartDing(ctx, cfg.Robots, a.robotHandler, a.logger.Logger) + } + if cfg.Robots.Wechat.Enabled && cfg.Robots.Wechat.BotToken != "" { + ctx, cancel := context.WithCancel(context.Background()) + a.wechatCancel = cancel + go robot.StartWechat(ctx, cfg.Robots, a.robotHandler, cfg.Version, a.logger.Logger) + } + if cfg.Robots.Telegram.Enabled && strings.TrimSpace(cfg.Robots.Telegram.BotToken) != "" { + ctx, cancel := context.WithCancel(context.Background()) + a.telegramCancel = cancel + go robot.StartTelegram(ctx, cfg.Robots, a.robotHandler, a.logger.Logger) + } + if cfg.Robots.Slack.Enabled && strings.TrimSpace(cfg.Robots.Slack.BotToken) != "" && strings.TrimSpace(cfg.Robots.Slack.AppToken) != "" { + ctx, cancel := context.WithCancel(context.Background()) + a.slackCancel = cancel + go robot.StartSlack(ctx, cfg.Robots, a.robotHandler, a.logger.Logger) + } + if cfg.Robots.Discord.Enabled && strings.TrimSpace(cfg.Robots.Discord.BotToken) != "" { + ctx, cancel := context.WithCancel(context.Background()) + a.discordCancel = cancel + go robot.StartDiscord(ctx, cfg.Robots, a.robotHandler, a.logger.Logger) + } + if cfg.Robots.QQ.Enabled && strings.TrimSpace(cfg.Robots.QQ.AppID) != "" && strings.TrimSpace(cfg.Robots.QQ.ClientSecret) != "" { + ctx, cancel := context.WithCancel(context.Background()) + a.qqCancel = cancel + go robot.StartQQ(ctx, cfg.Robots, a.robotHandler, a.logger.Logger) + } +} + +// RestartRobotConnections 重启钉钉/飞书/微信长连接,使前端应用配置后立即生效(实现 handler.RobotRestarter) +func (a *App) RestartRobotConnections() { + a.robotMu.Lock() + if a.dingCancel != nil { + a.dingCancel() + a.dingCancel = nil + } + if a.larkCancel != nil { + a.larkCancel() + a.larkCancel = nil + } + if a.wechatCancel != nil { + a.wechatCancel() + a.wechatCancel = nil + } + if a.telegramCancel != nil { + a.telegramCancel() + a.telegramCancel = nil + } + if a.slackCancel != nil { + a.slackCancel() + a.slackCancel = nil + } + if a.discordCancel != nil { + a.discordCancel() + a.discordCancel = nil + } + if a.qqCancel != nil { + a.qqCancel() + a.qqCancel = nil + } + a.robotMu.Unlock() + // 给旧 goroutine 一点时间退出 + time.Sleep(200 * time.Millisecond) + a.startRobotConnections() +} + +// setupRoutes 设置路由 +func setupRoutes( + router *gin.Engine, + authHandler *handler.AuthHandler, + agentHandler *handler.AgentHandler, + monitorHandler *handler.MonitorHandler, + notificationHandler *handler.NotificationHandler, + conversationHandler *handler.ConversationHandler, + robotHandler *handler.RobotHandler, + wechatRobotHandler *handler.WechatRobotHandler, + groupHandler *handler.GroupHandler, + configHandler *handler.ConfigHandler, + externalMCPHandler *handler.ExternalMCPHandler, + attackChainHandler *handler.AttackChainHandler, + app *App, // 传递 App 实例以便动态获取 knowledgeHandler + vulnerabilityHandler *handler.VulnerabilityHandler, + assetHandler *handler.AssetHandler, + projectHandler *handler.ProjectHandler, + workflowHandler *handler.WorkflowHandler, + webshellHandler *handler.WebShellHandler, + chatUploadsHandler *handler.ChatUploadsHandler, + roleHandler *handler.RoleHandler, + skillsHandler *handler.SkillsHandler, + markdownAgentsHandler *handler.MarkdownAgentsHandler, + fofaHandler *handler.FofaHandler, + terminalHandler *handler.TerminalHandler, + c2Handler *handler.C2Handler, + auditHandler *handler.AuditHandler, + auditSvc *audit.Service, + rbacHandler *handler.RBACHandler, + mcpServer *mcp.Server, + authManager *security.AuthManager, + openAPIHandler *handler.OpenAPIHandler, +) { + // API路由 + api := router.Group("/api") + + // 认证相关路由 + authRoutes := api.Group("/auth") + loginRL := security.NewRateLimiter(10, 1*time.Minute) + { + authRoutes.POST("/login", security.RateLimitMiddleware(loginRL), authHandler.Login) + authRoutes.POST("/logout", security.AuthMiddleware(authManager), authHandler.Logout) + authRoutes.POST("/change-password", security.AuthMiddleware(authManager), security.RequirePermission("auth:self"), authHandler.ChangePassword) + authRoutes.GET("/validate", security.AuthMiddleware(authManager), authHandler.Validate) + authRoutes.POST("/robot-binding-code", security.AuthMiddleware(authManager), security.RequirePermission("auth:self"), robotHandler.CreateRobotBindingCode) + authRoutes.GET("/robot-bindings", security.AuthMiddleware(authManager), security.RequirePermission("auth:self"), robotHandler.ListMyRobotBindings) + authRoutes.DELETE("/robot-bindings/:id", security.AuthMiddleware(authManager), security.RequirePermission("auth:self"), robotHandler.DeleteMyRobotBinding) + } + + // 机器人回调(无需登录,供企业微信/钉钉/飞书服务器调用) + // 添加速率限制:每个 IP 每分钟最多 60 次请求,防止滥用 + robotRL := security.NewRateLimiter(60, 1*time.Minute) + robotGroup := api.Group("/robot") + robotGroup.Use(security.RateLimitMiddleware(robotRL)) + { + robotGroup.GET("/wecom", robotHandler.HandleWecomGET) + robotGroup.POST("/wecom", robotHandler.HandleWecomPOST) + robotGroup.POST("/dingtalk", robotHandler.HandleDingtalkPOST) + robotGroup.POST("/lark", robotHandler.HandleLarkPOST) + } + + protected := api.Group("") + protected.Use(security.AuthMiddleware(authManager)) + protected.Use(security.RBACMiddlewareWithDenyHook(app.db, func(c *gin.Context, reason, permission string) { + if auditSvc != nil { + auditSvc.Record(c, audit.Entry{ + Level: "warn", Category: "rbac", Action: "access_denied", Result: "failure", + Message: "RBAC 拒绝访问", ResourceType: "route", ResourceID: c.FullPath(), + Detail: map[string]interface{}{"reason": reason, "permission": permission, "method": c.Request.Method}, + }) + } + })) + { + protected.GET("/rbac/me", rbacHandler.Me) + protected.GET("/rbac/metadata", rbacHandler.Metadata) + protected.GET("/rbac/users", rbacHandler.ListUsers) + protected.POST("/rbac/users", rbacHandler.CreateUser) + protected.PUT("/rbac/users/:id", rbacHandler.UpdateUser) + protected.DELETE("/rbac/users/:id", rbacHandler.DeleteUser) + protected.GET("/rbac/roles", rbacHandler.ListRoles) + protected.POST("/rbac/roles", rbacHandler.CreateRole) + protected.PUT("/rbac/roles/:id", rbacHandler.UpdateRole) + protected.DELETE("/rbac/roles/:id", rbacHandler.DeleteRole) + protected.GET("/rbac/resource-assignments", rbacHandler.ListResourceAssignments) + protected.GET("/rbac/resources", rbacHandler.ListAssignableResources) + protected.POST("/rbac/resource-assignments", rbacHandler.AssignResource) + protected.DELETE("/rbac/resource-assignments/:id", rbacHandler.DeleteResourceAssignment) + + // 机器人测试(需登录):POST /api/robot/test,body: {"platform":"dingtalk","user_id":"test","text":"帮助"},用于验证机器人逻辑 + protected.POST("/robot/test", robotHandler.HandleRobotTest) + + // 微信 iLink 扫码绑定(需登录) + protected.POST("/robot/wechat/qrcode", wechatRobotHandler.HandleWechatQRCode) + protected.GET("/robot/wechat/qrcode/status", wechatRobotHandler.HandleWechatQRCodeStatus) + protected.POST("/robot/wechat/qrcode/verify", wechatRobotHandler.HandleWechatVerifyCode) + protected.GET("/robot/wechat/status", wechatRobotHandler.HandleWechatStatus) + + // Eino ADK 单代理(ChatModelAgent + Runner;不依赖 multi_agent.enabled) + protected.POST("/eino-agent", agentHandler.EinoSingleAgentLoop) + protected.POST("/eino-agent/stream", agentHandler.EinoSingleAgentLoopStream) + protected.GET("/hitl/pending", agentHandler.ListHITLPending) + protected.GET("/hitl/logs", agentHandler.ListHITLLogs) + protected.DELETE("/hitl/logs", agentHandler.DeleteHITLLogs) + protected.GET("/hitl/logs/:id", agentHandler.GetHITLLog) + protected.POST("/hitl/decision", agentHandler.DecideHITLInterrupt) + protected.POST("/hitl/dismiss", agentHandler.DismissHITLInterrupt) + protected.GET("/hitl/config/:conversationId", agentHandler.GetHITLConversationConfig) + protected.PUT("/hitl/config", agentHandler.UpsertHITLConversationConfig) + protected.GET("/hitl/tool-whitelist", agentHandler.GetHITLGlobalToolWhitelist) + protected.PUT("/hitl/tool-whitelist", agentHandler.SetHITLGlobalToolWhitelist) + protected.POST("/hitl/tool-whitelist", agentHandler.MergeHITLGlobalToolWhitelist) + protected.GET("/hitl/default-reviewer", agentHandler.GetHITLDefaultReviewer) + protected.PUT("/hitl/default-reviewer", agentHandler.UpdateHITLDefaultReviewer) + protected.GET("/hitl/audit-strategy", agentHandler.GetHITLAuditStrategy) + protected.PUT("/hitl/audit-strategy", agentHandler.UpdateHITLAuditStrategy) + // Agent Loop 取消与任务列表 + protected.POST("/agent-loop/cancel", agentHandler.CancelAgentLoop) + protected.GET("/agent-loop/tasks", agentHandler.ListAgentTasks) + protected.GET("/agent-loop/task-events", agentHandler.SubscribeAgentTaskEvents) + protected.GET("/agent-loop/tasks/completed", agentHandler.ListCompletedTasks) + + // Eino DeepAgent 多代理(与单 Agent 并存,需 config.multi_agent.enabled) + // 多代理路由常注册;是否可用由运行时 h.config.MultiAgent.Enabled 决定(应用配置后无需重启) + protected.POST("/multi-agent", agentHandler.MultiAgentLoop) + protected.POST("/multi-agent/stream", agentHandler.MultiAgentLoopStream) + protected.GET("/multi-agent/markdown-agents", markdownAgentsHandler.ListMarkdownAgents) + protected.GET("/multi-agent/markdown-agents/:filename", markdownAgentsHandler.GetMarkdownAgent) + protected.POST("/multi-agent/markdown-agents", markdownAgentsHandler.CreateMarkdownAgent) + protected.PUT("/multi-agent/markdown-agents/:filename", markdownAgentsHandler.UpdateMarkdownAgent) + protected.DELETE("/multi-agent/markdown-agents/:filename", markdownAgentsHandler.DeleteMarkdownAgent) + + // 信息收集 - FOFA 查询(后端代理) + protected.POST("/fofa/search", fofaHandler.Search) + // 信息收集 - 自然语言解析为 FOFA 语法(需人工确认后再查询) + protected.POST("/fofa/parse", fofaHandler.ParseNaturalLanguage) + + // 资产管理 + protected.GET("/assets", assetHandler.List) + protected.GET("/assets/selection", assetHandler.Selection) + protected.GET("/assets/stats", assetHandler.Stats) + protected.POST("/assets/import", assetHandler.Import) + protected.POST("/assets/scan-links", assetHandler.RecordScans) + protected.PUT("/assets/bulk", assetHandler.BulkUpdate) + protected.PUT("/assets/project-binding", assetHandler.UpdateProjectBinding) + protected.POST("/assets/batch-delete", assetHandler.BatchDelete) + protected.POST("/assets/merge", security.RequirePermission("asset:write"), assetHandler.Merge) + protected.PUT("/assets/:id", assetHandler.Update) + protected.DELETE("/assets/:id", assetHandler.Delete) + + // 批量任务管理 + protected.POST("/batch-tasks", agentHandler.CreateBatchQueue) + protected.GET("/batch-tasks", agentHandler.ListBatchQueues) + protected.GET("/batch-tasks/:queueId", agentHandler.GetBatchQueue) + protected.POST("/batch-tasks/:queueId/start", agentHandler.StartBatchQueue) + protected.POST("/batch-tasks/:queueId/rerun", agentHandler.RerunBatchQueue) + protected.POST("/batch-tasks/:queueId/pause", agentHandler.PauseBatchQueue) + protected.PUT("/batch-tasks/:queueId/metadata", agentHandler.UpdateBatchQueueMetadata) + protected.PUT("/batch-tasks/:queueId/schedule", agentHandler.UpdateBatchQueueSchedule) + protected.PUT("/batch-tasks/:queueId/schedule-enabled", agentHandler.SetBatchQueueScheduleEnabled) + protected.DELETE("/batch-tasks/:queueId", agentHandler.DeleteBatchQueue) + protected.PUT("/batch-tasks/:queueId/tasks/:taskId", agentHandler.UpdateBatchTask) + protected.POST("/batch-tasks/:queueId/tasks/:taskId/run", agentHandler.RunSingleBatchTask) + protected.POST("/batch-tasks/:queueId/tasks", agentHandler.AddBatchTask) + protected.DELETE("/batch-tasks/:queueId/tasks/:taskId", agentHandler.DeleteBatchTask) + + // 对话历史 + protected.POST("/conversations", conversationHandler.CreateConversation) + protected.GET("/conversations", conversationHandler.ListConversations) + protected.GET("/conversations/:id", conversationHandler.GetConversation) + protected.GET("/conversations/:id/plan-tasks", conversationHandler.GetConversationPlanTasks) + protected.GET("/messages/:id/process-details", conversationHandler.GetMessageProcessDetails) + protected.GET("/process-details/:id", conversationHandler.GetProcessDetail) + protected.PUT("/conversations/:id", conversationHandler.UpdateConversation) + protected.PUT("/conversations/:id/project", conversationHandler.SetConversationProject) + protected.DELETE("/conversations/:id", conversationHandler.DeleteConversation) + protected.POST("/conversations/:id/delete-turn", conversationHandler.DeleteConversationTurn) + protected.PUT("/conversations/:id/pinned", groupHandler.UpdateConversationPinned) + + // 对话分组 + protected.POST("/groups", groupHandler.CreateGroup) + protected.GET("/groups", groupHandler.ListGroups) + protected.GET("/groups/:id", groupHandler.GetGroup) + protected.PUT("/groups/:id", groupHandler.UpdateGroup) + protected.DELETE("/groups/:id", groupHandler.DeleteGroup) + protected.PUT("/groups/:id/pinned", groupHandler.UpdateGroupPinned) + protected.GET("/groups/:id/conversations", groupHandler.GetGroupConversations) + protected.GET("/groups/mappings", groupHandler.GetAllMappings) + protected.POST("/groups/conversations", groupHandler.AddConversationToGroup) + protected.DELETE("/groups/:id/conversations/:conversationId", groupHandler.RemoveConversationFromGroup) + protected.PUT("/groups/:id/conversations/:conversationId/pinned", groupHandler.UpdateConversationPinnedInGroup) + + // 监控 + protected.GET("/monitor", monitorHandler.Monitor) + protected.GET("/monitor/execution/:id", monitorHandler.GetExecution) + protected.POST("/monitor/execution/:id/cancel", monitorHandler.CancelExecution) + protected.POST("/monitor/executions/names", monitorHandler.BatchGetToolNames) + protected.DELETE("/monitor/execution/:id", monitorHandler.DeleteExecution) + protected.DELETE("/monitor/executions", monitorHandler.DeleteExecutions) + protected.GET("/monitor/stats", monitorHandler.GetStats) + protected.GET("/monitor/calls-timeline", monitorHandler.GetCallsTimeline) + protected.GET("/notifications/summary", notificationHandler.GetSummary) + protected.POST("/notifications/read", notificationHandler.MarkRead) + + // 配置管理 + protected.GET("/config", configHandler.GetConfig) + protected.GET("/config/tools", configHandler.GetTools) + protected.GET("/config/tools/:name/schema", configHandler.GetToolSchema) + protected.PUT("/config", configHandler.UpdateConfig) + protected.POST("/config/apply", configHandler.ApplyConfig) + protected.POST("/config/test-openai", configHandler.TestOpenAI) + protected.POST("/config/test-vision", configHandler.TestVision) + protected.POST("/config/list-models", configHandler.ListModels) + + // 系统设置 - 终端(执行命令,提高运维效率) + protected.POST("/terminal/run", terminalHandler.RunCommand) + protected.POST("/terminal/run/stream", terminalHandler.RunCommandStream) + protected.GET("/terminal/ws", terminalHandler.RunCommandWS) + + // 平台审计日志 + protected.GET("/audit/meta", auditHandler.Meta) + protected.GET("/audit/summary", auditHandler.Summary) + protected.GET("/audit/logs", auditHandler.ListLogs) + protected.GET("/audit/logs/export", auditHandler.ExportLogs) + protected.GET("/audit/logs/:id", auditHandler.GetLog) + + // 外部MCP管理 + protected.GET("/external-mcp", externalMCPHandler.GetExternalMCPs) + protected.GET("/external-mcp/stats", externalMCPHandler.GetExternalMCPStats) + protected.GET("/external-mcp/:name", externalMCPHandler.GetExternalMCP) + protected.PUT("/external-mcp/:name", externalMCPHandler.AddOrUpdateExternalMCP) + protected.DELETE("/external-mcp/:name", externalMCPHandler.DeleteExternalMCP) + protected.POST("/external-mcp/:name/start", externalMCPHandler.StartExternalMCP) + protected.POST("/external-mcp/:name/stop", externalMCPHandler.StopExternalMCP) + + // 攻击链可视化 + protected.GET("/attack-chain/:conversationId", attackChainHandler.GetAttackChain) + protected.POST("/attack-chain/:conversationId/regenerate", attackChainHandler.RegenerateAttackChain) + + // 知识库管理(始终注册路由,通过 App 实例动态获取 handler) + knowledgeRoutes := protected.Group("/knowledge") + { + knowledgeRoutes.GET("/categories", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "categories": []string{}, + "enabled": false, + "message": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.GetCategories(c) + }) + knowledgeRoutes.GET("/items", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "items": []interface{}{}, + "enabled": false, + "message": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.GetItems(c) + }) + knowledgeRoutes.GET("/items/:id", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "enabled": false, + "message": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.GetItem(c) + }) + knowledgeRoutes.POST("/items", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "enabled": false, + "error": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.CreateItem(c) + }) + knowledgeRoutes.PUT("/items/:id", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "enabled": false, + "error": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.UpdateItem(c) + }) + knowledgeRoutes.DELETE("/items/:id", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "enabled": false, + "error": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.DeleteItem(c) + }) + knowledgeRoutes.GET("/index-status", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "enabled": false, + "total_items": 0, + "indexed_items": 0, + "progress_percent": 0, + "is_complete": false, + "message": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.GetIndexStatus(c) + }) + knowledgeRoutes.POST("/index", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "enabled": false, + "error": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.StartIndex(c) + }) + knowledgeRoutes.POST("/scan", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "enabled": false, + "error": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.ScanKnowledgeBase(c) + }) + knowledgeRoutes.GET("/retrieval-logs", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "logs": []interface{}{}, + "enabled": false, + "message": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.GetRetrievalLogs(c) + }) + knowledgeRoutes.DELETE("/retrieval-logs/:id", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "enabled": false, + "error": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.DeleteRetrievalLog(c) + }) + knowledgeRoutes.POST("/search", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "results": []interface{}{}, + "enabled": false, + "message": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.Search(c) + }) + knowledgeRoutes.GET("/stats", func(c *gin.Context) { + if app.knowledgeHandler == nil { + c.JSON(http.StatusOK, gin.H{ + "enabled": false, + "total_categories": 0, + "total_items": 0, + "message": "知识库功能未启用,请前往系统设置启用知识检索功能", + }) + return + } + app.knowledgeHandler.GetStats(c) + }) + } + + // 漏洞管理 + protected.GET("/vulnerabilities", vulnerabilityHandler.ListVulnerabilities) + protected.GET("/vulnerabilities/export", vulnerabilityHandler.ExportVulnerabilities) + protected.DELETE("/vulnerabilities/batch", vulnerabilityHandler.BatchDeleteVulnerabilities) + protected.GET("/vulnerabilities/filter-options", vulnerabilityHandler.GetVulnerabilityFilterOptions) + protected.GET("/vulnerabilities/stats", vulnerabilityHandler.GetVulnerabilityStats) + protected.GET("/vulnerability-alerts/subscription", vulnerabilityHandler.GetMyAlertSubscription) + protected.PUT("/vulnerability-alerts/subscription", vulnerabilityHandler.UpdateMyAlertSubscription) + protected.GET("/vulnerabilities/:id", vulnerabilityHandler.GetVulnerability) + protected.POST("/vulnerabilities", vulnerabilityHandler.CreateVulnerability) + protected.PUT("/vulnerabilities/:id", vulnerabilityHandler.UpdateVulnerability) + protected.DELETE("/vulnerabilities/:id", vulnerabilityHandler.DeleteVulnerability) + + // 项目管理与事实黑板 + protected.GET("/projects/dashboard-summary", projectHandler.GetDashboardSummary) + protected.GET("/projects", projectHandler.ListProjects) + protected.POST("/projects", projectHandler.CreateProject) + protected.GET("/projects/:id/stats", projectHandler.GetProjectStats) + protected.GET("/projects/:id/conversations", projectHandler.ListProjectConversations) + protected.GET("/projects/:id", projectHandler.GetProject) + protected.PUT("/projects/:id", projectHandler.UpdateProject) + protected.DELETE("/projects/:id", projectHandler.DeleteProject) + protected.GET("/projects/:id/fact-graph", projectHandler.GetFactGraph) + protected.GET("/projects/:id/fact-edges", projectHandler.ListFactEdges) + protected.POST("/projects/:id/fact-edges", projectHandler.CreateFactEdge) + protected.DELETE("/projects/:id/fact-edges/:edgeId", projectHandler.DeleteFactEdge) + protected.POST("/projects/:id/promote-attack-chain/:conversationId", projectHandler.PromoteAttackChain) + protected.GET("/projects/:id/facts", projectHandler.ListFacts) + protected.POST("/projects/:id/facts", projectHandler.CreateFact) + protected.PUT("/projects/:id/facts/:factId", projectHandler.UpdateFact) + protected.DELETE("/projects/:id/facts/:factId", projectHandler.DeleteFact) + protected.POST("/projects/:id/facts/deprecate", projectHandler.DeprecateFact) + protected.POST("/projects/:id/facts/restore", projectHandler.RestoreFact) + + // WebShell 管理(代理执行 + 连接配置存 SQLite) + protected.GET("/webshell/connections", webshellHandler.ListConnections) + protected.POST("/webshell/connections", webshellHandler.CreateConnection) + protected.GET("/webshell/connections/:id/ai-history", webshellHandler.GetAIHistory) + protected.GET("/webshell/connections/:id/ai-conversations", webshellHandler.ListAIConversations) + protected.GET("/webshell/connections/:id/state", webshellHandler.GetConnectionState) + protected.PUT("/webshell/connections/:id", webshellHandler.UpdateConnection) + protected.PUT("/webshell/connections/:id/state", webshellHandler.SaveConnectionState) + protected.DELETE("/webshell/connections/:id", webshellHandler.DeleteConnection) + protected.POST("/webshell/exec", webshellHandler.Exec) + protected.POST("/webshell/file", webshellHandler.FileOp) + + // C2 管理(未启用时返回 503,避免 Handler 空指针) + c2Routes := protected.Group("/c2") + c2Routes.Use(func(c *gin.Context) { + if app.c2Manager == nil { + c.AbortWithStatusJSON(http.StatusServiceUnavailable, gin.H{ + "error": "c2_disabled", + "message": "C2 功能已在系统设置中关闭", + "enabled": false, + }) + return + } + c.Next() + }) + c2Routes.GET("/listeners", c2Handler.ListListeners) + c2Routes.POST("/listeners", c2Handler.CreateListener) + c2Routes.GET("/listeners/:id", c2Handler.GetListener) + c2Routes.PUT("/listeners/:id", c2Handler.UpdateListener) + c2Routes.DELETE("/listeners/:id", c2Handler.DeleteListener) + c2Routes.POST("/listeners/:id/start", c2Handler.StartListener) + c2Routes.POST("/listeners/:id/stop", c2Handler.StopListener) + c2Routes.GET("/sessions", c2Handler.ListSessions) + c2Routes.DELETE("/sessions", c2Handler.DeleteSessions) + c2Routes.GET("/sessions/:id", c2Handler.GetSession) + c2Routes.DELETE("/sessions/:id", c2Handler.DeleteSession) + c2Routes.PUT("/sessions/:id/sleep", c2Handler.SetSessionSleep) + c2Routes.PUT("/sessions/:id/note", c2Handler.SetSessionNote) + c2Routes.GET("/tasks", c2Handler.ListTasks) + c2Routes.DELETE("/tasks", c2Handler.DeleteTasks) + c2Routes.GET("/tasks/:id", c2Handler.GetTask) + c2Routes.POST("/tasks", c2Handler.CreateTask) + c2Routes.POST("/tasks/:id/cancel", c2Handler.CancelTask) + c2Routes.GET("/tasks/:id/wait", c2Handler.WaitTask) + c2Routes.POST("/sessions/:id/tasks", c2Handler.CreateTask) + c2Routes.POST("/payloads/oneliner", c2Handler.PayloadOneliner) + c2Routes.POST("/payloads/build", c2Handler.PayloadBuild) + c2Routes.GET("/payloads/:id/download", c2Handler.PayloadDownload) + c2Routes.GET("/events", c2Handler.ListEvents) + c2Routes.DELETE("/events", c2Handler.DeleteEvents) + c2Routes.GET("/events/stream", c2Handler.EventStream) + c2Routes.POST("/files/upload", c2Handler.UploadFileForImplant) + c2Routes.GET("/files", c2Handler.ListFiles) + c2Routes.GET("/tasks/:id/result-file", c2Handler.DownloadResultFile) + c2Routes.GET("/profiles", c2Handler.ListProfiles) + c2Routes.GET("/profiles/:id", c2Handler.GetProfile) + c2Routes.POST("/profiles", c2Handler.CreateProfile) + c2Routes.PUT("/profiles/:id", c2Handler.UpdateProfile) + c2Routes.DELETE("/profiles/:id", c2Handler.DeleteProfile) + + // 对话附件(chat_uploads)管理 + protected.GET("/chat-uploads", chatUploadsHandler.List) + protected.GET("/chat-uploads/export", chatUploadsHandler.Export) + protected.GET("/chat-uploads/download", chatUploadsHandler.Download) + protected.GET("/chat-uploads/path", chatUploadsHandler.ResolvePath) + protected.GET("/chat-uploads/content", chatUploadsHandler.GetContent) + protected.POST("/chat-uploads", chatUploadsHandler.Upload) + protected.POST("/chat-uploads/mkdir", chatUploadsHandler.Mkdir) + protected.DELETE("/chat-uploads", chatUploadsHandler.Delete) + protected.PUT("/chat-uploads/rename", chatUploadsHandler.Rename) + protected.PUT("/chat-uploads/content", chatUploadsHandler.PutContent) + + // 角色管理 + protected.GET("/roles", roleHandler.GetRoles) + protected.GET("/roles/:name", roleHandler.GetRole) + protected.POST("/roles", roleHandler.CreateRole) + protected.PUT("/roles/:name", roleHandler.UpdateRole) + protected.DELETE("/roles/:name", roleHandler.DeleteRole) + + // 工作流定义(图结构固定,业务字段保存在 graph_json 中) + protected.GET("/workflows/runs/pending", workflowHandler.ListPendingRuns) + protected.GET("/workflows/runs/:runId/replay", workflowHandler.ReplayRun) + protected.GET("/workflows/runs/:runId", workflowHandler.GetRun) + protected.POST("/workflows/runs/:runId/resume", workflowHandler.ResumeRun) + protected.POST("/workflows/validate", workflowHandler.Validate) + protected.POST("/workflows/dry-run", workflowHandler.DryRun) + protected.POST("/workflows/generate-draft", workflowHandler.GenerateDraft) + protected.GET("/workflows/:id/package", workflowHandler.ExportPackage) + protected.POST("/workflow-package-inspections", workflowHandler.CreatePackageInspection) + protected.GET("/workflow-package-inspections/:inspectionId", workflowHandler.GetPackageInspection) + protected.POST("/workflow-package-imports", workflowHandler.ApplyPackageImport) + protected.GET("/workflow-package-imports/:importId", workflowHandler.GetPackageImport) + protected.GET("/workflows", workflowHandler.List) + protected.GET("/workflows/:id", workflowHandler.Get) + protected.POST("/workflows", workflowHandler.Create) + protected.PUT("/workflows/:id", workflowHandler.Update) + protected.DELETE("/workflows/:id", workflowHandler.Delete) + + // Skills管理(具体路径需注册在 /skills/:name 之前) + protected.GET("/skills", skillsHandler.GetSkills) + protected.GET("/skills/stats", skillsHandler.GetSkillStats) + protected.DELETE("/skills/stats", skillsHandler.ClearSkillStats) + protected.GET("/skills/:name/files", skillsHandler.ListSkillPackageFiles) + protected.GET("/skills/:name/file", skillsHandler.GetSkillPackageFile) + protected.PUT("/skills/:name/file", skillsHandler.PutSkillPackageFile) + protected.GET("/skills/:name/bound-roles", skillsHandler.GetSkillBoundRoles) + protected.POST("/skills", skillsHandler.CreateSkill) + protected.PUT("/skills/:name", skillsHandler.UpdateSkill) + protected.DELETE("/skills/:name", skillsHandler.DeleteSkill) + protected.DELETE("/skills/:name/stats", skillsHandler.ClearSkillStatsByName) + protected.GET("/skills/:name", skillsHandler.GetSkill) + + // MCP端点 + protected.POST("/mcp", func(c *gin.Context) { + mcpServer.HandleHTTP(c.Writer, c.Request) + }) + + // OpenAPI结果聚合端点(可选,用于获取对话的完整结果) + protected.GET("/conversations/:id/results", openAPIHandler.GetConversationResults) + } + + // OpenAPI规范(需要认证,避免暴露API结构信息) + protected.GET("/openapi/spec", openAPIHandler.GetOpenAPISpec) + + // API文档页面(公开访问,但需要登录后才能使用API) + router.GET("/api-docs", func(c *gin.Context) { + c.HTML(http.StatusOK, "api-docs.html", nil) + }) + + // 静态文件 + router.Static("/static", "./web/static") + router.LoadHTMLGlob("web/templates/*") + + // 前端页面 + router.GET("/", func(c *gin.Context) { + version := app.config.Version + if version == "" { + version = "v1.0.0" + } + c.HTML(http.StatusOK, "index.html", gin.H{"Version": version}) + }) +} + +// registerWebshellTools 注册 WebShell 相关 MCP 工具,供 AI 助手在指定连接上执行命令与文件操作 +func registerWebshellTools(mcpServer *mcp.Server, db *database.DB, webshellHandler *handler.WebShellHandler, logger *zap.Logger) { + if db == nil || webshellHandler == nil { + logger.Warn("跳过 WebShell 工具注册:db 或 webshellHandler 为空") + return + } + + // webshell_exec + execTool := mcp.Tool{ + Name: builtin.ToolWebshellExec, + Description: "在指定的 WebShell 连接上执行一条系统命令,返回命令的标准输出。connection_id 由用户在 AI 助手上下文中选定。", + ShortDescription: "在 WebShell 连接上执行命令", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "connection_id": map[string]interface{}{ + "type": "string", + "description": "WebShell 连接 ID(如 ws_xxx)", + }, + "command": map[string]interface{}{ + "type": "string", + "description": "要执行的系统命令", + }, + }, + "required": []string{"connection_id", "command"}, + }, + } + execHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + cid, _ := args["connection_id"].(string) + cmd, _ := args["command"].(string) + if cid == "" || cmd == "" { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "connection_id 和 command 均为必填"}}, IsError: true}, nil + } + conn, err := db.GetWebshellConnection(cid) + if err != nil || conn == nil { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "未找到该 WebShell 连接或查询失败"}}, IsError: true}, nil + } + output, ok, errMsg := webshellHandler.ExecWithConnection(conn, cmd) + if errMsg != "" { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: errMsg}}, IsError: true}, nil + } + if !ok { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "HTTP 非 200,输出:\n" + output}}, IsError: false}, nil + } + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: output}}, IsError: false}, nil + } + mcpServer.RegisterTool(execTool, execHandler) + + // webshell_file_list + listTool := mcp.Tool{ + Name: builtin.ToolWebshellFileList, + Description: "在指定 WebShell 连接上列出目录内容。path 默认为当前目录(.)。", + ShortDescription: "在 WebShell 上列出目录", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "connection_id": map[string]interface{}{"type": "string", "description": "WebShell 连接 ID"}, + "path": map[string]interface{}{"type": "string", "description": "目录路径,默认 ."}, + }, + "required": []string{"connection_id"}, + }, + } + listHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + cid, _ := args["connection_id"].(string) + path, _ := args["path"].(string) + if cid == "" { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "connection_id 必填"}}, IsError: true}, nil + } + conn, err := db.GetWebshellConnection(cid) + if err != nil || conn == nil { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "未找到该 WebShell 连接"}}, IsError: true}, nil + } + output, ok, errMsg := webshellHandler.FileOpWithConnection(conn, "list", path, "", "") + if errMsg != "" { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: errMsg}}, IsError: true}, nil + } + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: output}}, IsError: !ok}, nil + } + mcpServer.RegisterTool(listTool, listHandler) + + // webshell_file_read + readTool := mcp.Tool{ + Name: builtin.ToolWebshellFileRead, + Description: "在指定 WebShell 连接上读取文件内容。", + ShortDescription: "在 WebShell 上读取文件", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "connection_id": map[string]interface{}{"type": "string", "description": "WebShell 连接 ID"}, + "path": map[string]interface{}{"type": "string", "description": "文件路径"}, + }, + "required": []string{"connection_id", "path"}, + }, + } + readHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + cid, _ := args["connection_id"].(string) + path, _ := args["path"].(string) + if cid == "" || path == "" { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "connection_id 和 path 必填"}}, IsError: true}, nil + } + conn, err := db.GetWebshellConnection(cid) + if err != nil || conn == nil { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "未找到该 WebShell 连接"}}, IsError: true}, nil + } + output, ok, errMsg := webshellHandler.FileOpWithConnection(conn, "read", path, "", "") + if errMsg != "" { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: errMsg}}, IsError: true}, nil + } + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: output}}, IsError: !ok}, nil + } + mcpServer.RegisterTool(readTool, readHandler) + + // webshell_file_write + writeTool := mcp.Tool{ + Name: builtin.ToolWebshellFileWrite, + Description: "在指定 WebShell 连接上写入文件内容(会覆盖已有文件)。", + ShortDescription: "在 WebShell 上写入文件", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "connection_id": map[string]interface{}{"type": "string", "description": "WebShell 连接 ID"}, + "path": map[string]interface{}{"type": "string", "description": "文件路径"}, + "content": map[string]interface{}{"type": "string", "description": "要写入的内容"}, + }, + "required": []string{"connection_id", "path", "content"}, + }, + } + writeHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + cid, _ := args["connection_id"].(string) + path, _ := args["path"].(string) + content, _ := args["content"].(string) + if cid == "" || path == "" { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "connection_id 和 path 必填"}}, IsError: true}, nil + } + conn, err := db.GetWebshellConnection(cid) + if err != nil || conn == nil { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "未找到该 WebShell 连接"}}, IsError: true}, nil + } + output, ok, errMsg := webshellHandler.FileOpWithConnection(conn, "write", path, content, "") + if errMsg != "" { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: errMsg}}, IsError: true}, nil + } + if !ok { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "写入可能失败,输出:\n" + output}}, IsError: false}, nil + } + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "写入成功\n" + output}}, IsError: false}, nil + } + mcpServer.RegisterTool(writeTool, writeHandler) + + logger.Debug("WebShell 工具注册成功") +} + +// registerWebshellManagementTools 注册 WebShell 连接管理 MCP 工具 +func registerWebshellManagementTools(mcpServer *mcp.Server, db *database.DB, webshellHandler *handler.WebShellHandler, logger *zap.Logger) { + if db == nil { + logger.Warn("跳过 WebShell 管理工具注册:db 为空") + return + } + projectIDFromToolArgs := func(ctx context.Context, args map[string]interface{}) string { + projectID, _ := args["project_id"].(string) + projectID = strings.TrimSpace(projectID) + if projectID == "" { + projectID = strings.TrimSpace(mcp.MCPProjectIDFromContext(ctx)) + } + return projectID + } + explicitProjectIDFromToolArgs := func(args map[string]interface{}) string { + projectID, _ := args["project_id"].(string) + return strings.TrimSpace(projectID) + } + authorizeWebshellToolProject := func(principal authctx.Principal, permission, projectID string) *mcp.ToolResult { + projectID = strings.TrimSpace(projectID) + if projectID == "" { + return nil + } + if projectID == database.ProjectFilterUnbound { + return nil + } + if !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), "project", projectID) { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "无权访问项目: " + projectID}}, + IsError: true, + } + } + return nil + } + + // manage_webshell_list - 列出所有 webshell 连接 + listTool := mcp.Tool{ + Name: builtin.ToolManageWebshellList, + Description: "列出已保存的 WebShell 连接,返回连接ID、URL、类型、所属项目、备注等信息。默认按当前对话项目边界过滤:项目对话看本项目,未绑定项目的对话看未绑定连接;显式传 project_id 时按指定项目过滤。", + ShortDescription: "列出所有 WebShell 连接", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "project_id": map[string]interface{}{ + "type": "string", + "description": "项目 ID;不填时在项目会话中默认使用当前项目。", + }, + }, + }, + } + listHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + connections := []database.WebShellConnection{} + var err error + if principal, ok := authctx.PrincipalFromContext(ctx); ok { + projectID := explicitProjectIDFromToolArgs(args) + if projectID == "" { + projectID = mcpEffectiveProjectFilter(ctx, db) + } + if result := authorizeWebshellToolProject(principal, "webshell:read", projectID); result != nil { + return result, nil + } + connections, err = db.ListWebshellConnectionsForAccess(principal.UserID, principal.ScopeFor("webshell:read"), projectID) + } else { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "缺少认证身份"}}, IsError: true}, nil + } + if err != nil { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "获取连接列表失败: " + err.Error()}}, + IsError: true, + }, nil + } + if len(connections) == 0 { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "暂无 WebShell 连接"}}, + IsError: false, + }, nil + } + var sb strings.Builder + sb.WriteString(fmt.Sprintf("找到 %d 个 WebShell 连接:\n\n", len(connections))) + for _, conn := range connections { + sb.WriteString(fmt.Sprintf("ID: %s\n", conn.ID)) + sb.WriteString(fmt.Sprintf(" URL: %s\n", conn.URL)) + sb.WriteString(fmt.Sprintf(" 类型: %s\n", conn.Type)) + sb.WriteString(fmt.Sprintf(" 请求方式: %s\n", conn.Method)) + sb.WriteString(fmt.Sprintf(" 命令参数: %s\n", conn.CmdParam)) + if conn.ProjectID != "" { + sb.WriteString(fmt.Sprintf(" 项目ID: %s\n", conn.ProjectID)) + } else { + sb.WriteString(" 项目: 未绑定\n") + } + if conn.Remark != "" { + sb.WriteString(fmt.Sprintf(" 备注: %s\n", conn.Remark)) + } + sb.WriteString(fmt.Sprintf(" 创建时间: %s\n", conn.CreatedAt.Format("2006-01-02 15:04:05"))) + sb.WriteString("\n") + } + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: sb.String()}}, + IsError: false, + }, nil + } + mcpServer.RegisterTool(listTool, listHandler) + + // manage_webshell_add - 添加新的 webshell 连接 + addTool := mcp.Tool{ + Name: builtin.ToolManageWebshellAdd, + Description: "添加新的 WebShell 连接到管理系统。支持 PHP、ASP、ASPX、JSP 等类型的一句话木马。", + ShortDescription: "添加 WebShell 连接", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "url": map[string]interface{}{ + "type": "string", + "description": "Shell 地址,如 http://target.com/shell.php(必填)", + }, + "password": map[string]interface{}{ + "type": "string", + "description": "连接密码/密钥,如冰蝎/蚁剑的连接密码", + }, + "type": map[string]interface{}{ + "type": "string", + "description": "Shell 类型:php、asp、aspx、jsp,默认为 php", + "enum": []string{"php", "asp", "aspx", "jsp"}, + }, + "method": map[string]interface{}{ + "type": "string", + "description": "请求方式:GET 或 POST,默认为 POST", + "enum": []string{"GET", "POST"}, + }, + "cmd_param": map[string]interface{}{ + "type": "string", + "description": "命令参数名,不填默认为 cmd", + }, + "remark": map[string]interface{}{ + "type": "string", + "description": "备注,便于识别的备注名", + }, + }, + "required": []string{"url"}, + }, + } + addHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + urlStr, _ := args["url"].(string) + if urlStr == "" { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "错误: url 参数必填"}}, + IsError: true, + }, nil + } + + password, _ := args["password"].(string) + shellType, _ := args["type"].(string) + if shellType == "" { + shellType = "php" + } + method, _ := args["method"].(string) + if method == "" { + method = "post" + } + cmdParam, _ := args["cmd_param"].(string) + if cmdParam == "" { + cmdParam = "cmd" + } + remark, _ := args["remark"].(string) + principal, ok := authctx.PrincipalFromContext(ctx) + if !ok { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "缺少认证身份"}}, IsError: true}, nil + } + projectID := projectIDFromToolArgs(ctx, args) + if result := authorizeWebshellToolProject(principal, "webshell:write", projectID); result != nil { + return result, nil + } + + // 生成连接ID + connID := "ws_" + strings.ReplaceAll(uuid.New().String(), "-", "")[:12] + conn := &database.WebShellConnection{ + ID: connID, + URL: urlStr, + Password: password, + Type: strings.ToLower(shellType), + Method: strings.ToLower(method), + CmdParam: cmdParam, + Remark: remark, + ProjectID: projectID, + CreatedAt: time.Now(), + } + + if err := db.CreateWebshellConnection(conn); err != nil { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "添加 WebShell 连接失败: " + err.Error()}}, + IsError: true, + }, nil + } + _ = db.SetResourceOwner("webshell", conn.ID, principal.UserID) + _ = db.AssignResourceToUser(principal.UserID, "webshell", conn.ID) + projectLine := "项目: 未绑定" + if conn.ProjectID != "" { + projectLine = "项目ID: " + conn.ProjectID + } + + return &mcp.ToolResult{ + Content: []mcp.Content{{ + Type: "text", + Text: fmt.Sprintf("WebShell 连接添加成功!\n\n连接ID: %s\nURL: %s\n类型: %s\n请求方式: %s\n命令参数: %s\n%s", conn.ID, conn.URL, conn.Type, conn.Method, conn.CmdParam, projectLine), + }}, + IsError: false, + }, nil + } + mcpServer.RegisterTool(addTool, addHandler) + + // manage_webshell_update - 更新 webshell 连接 + updateTool := mcp.Tool{ + Name: builtin.ToolManageWebshellUpdate, + Description: "更新已存在的 WebShell 连接信息。", + ShortDescription: "更新 WebShell 连接", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "connection_id": map[string]interface{}{ + "type": "string", + "description": "要更新的 WebShell 连接 ID(必填)", + }, + "url": map[string]interface{}{ + "type": "string", + "description": "新的 Shell 地址", + }, + "password": map[string]interface{}{ + "type": "string", + "description": "新的连接密码/密钥", + }, + "type": map[string]interface{}{ + "type": "string", + "description": "新的 Shell 类型:php、asp、aspx、jsp", + "enum": []string{"php", "asp", "aspx", "jsp"}, + }, + "method": map[string]interface{}{ + "type": "string", + "description": "新的请求方式:GET 或 POST", + "enum": []string{"GET", "POST"}, + }, + "cmd_param": map[string]interface{}{ + "type": "string", + "description": "新的命令参数名", + }, + "remark": map[string]interface{}{ + "type": "string", + "description": "新的备注", + }, + "project_id": map[string]interface{}{ + "type": "string", + "description": "新的所属项目 ID;传空字符串可取消绑定。", + }, + }, + "required": []string{"connection_id"}, + }, + } + updateHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + connID, _ := args["connection_id"].(string) + if connID == "" { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "错误: connection_id 参数必填"}}, + IsError: true, + }, nil + } + + // 获取现有连接 + existing, err := db.GetWebshellConnection(connID) + if err != nil || existing == nil { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "未找到指定的 WebShell 连接: " + connID}}, + IsError: true, + }, nil + } + + // 更新字段(如果提供了新值) + if urlStr, ok := args["url"].(string); ok && urlStr != "" { + existing.URL = urlStr + } + if password, ok := args["password"].(string); ok { + existing.Password = password + } + if shellType, ok := args["type"].(string); ok && shellType != "" { + existing.Type = strings.ToLower(shellType) + } + if method, ok := args["method"].(string); ok && method != "" { + existing.Method = strings.ToLower(method) + } + if cmdParam, ok := args["cmd_param"].(string); ok && cmdParam != "" { + existing.CmdParam = cmdParam + } + if remark, ok := args["remark"].(string); ok { + existing.Remark = remark + } + if projectID, ok := args["project_id"].(string); ok { + projectID = strings.TrimSpace(projectID) + if projectID != "" { + principal, ok := authctx.PrincipalFromContext(ctx) + if !ok { + return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "缺少认证身份"}}, IsError: true}, nil + } + if result := authorizeWebshellToolProject(principal, "webshell:write", projectID); result != nil { + return result, nil + } + } + existing.ProjectID = projectID + } + + if err := db.UpdateWebshellConnection(existing); err != nil { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "更新 WebShell 连接失败: " + err.Error()}}, + IsError: true, + }, nil + } + + return &mcp.ToolResult{ + Content: []mcp.Content{{ + Type: "text", + Text: fmt.Sprintf("WebShell 连接更新成功!\n\n连接ID: %s\nURL: %s\n类型: %s\n请求方式: %s\n命令参数: %s\n项目ID: %s\n备注: %s", existing.ID, existing.URL, existing.Type, existing.Method, existing.CmdParam, existing.ProjectID, existing.Remark), + }}, + IsError: false, + }, nil + } + mcpServer.RegisterTool(updateTool, updateHandler) + + // manage_webshell_delete - 删除 webshell 连接 + deleteTool := mcp.Tool{ + Name: builtin.ToolManageWebshellDelete, + Description: "删除指定的 WebShell 连接。", + ShortDescription: "删除 WebShell 连接", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "connection_id": map[string]interface{}{ + "type": "string", + "description": "要删除的 WebShell 连接 ID(必填)", + }, + }, + "required": []string{"connection_id"}, + }, + } + deleteHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + connID, _ := args["connection_id"].(string) + if connID == "" { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "错误: connection_id 参数必填"}}, + IsError: true, + }, nil + } + + if err := db.DeleteWebshellConnection(connID); err != nil { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "删除 WebShell 连接失败: " + err.Error()}}, + IsError: true, + }, nil + } + + return &mcp.ToolResult{ + Content: []mcp.Content{{ + Type: "text", + Text: fmt.Sprintf("WebShell 连接 %s 已成功删除", connID), + }}, + IsError: false, + }, nil + } + mcpServer.RegisterTool(deleteTool, deleteHandler) + + // manage_webshell_test - 测试 webshell 连接 + testTool := mcp.Tool{ + Name: builtin.ToolManageWebshellTest, + Description: "测试指定的 WebShell 连接是否可用,会尝试执行一个简单的命令(如 whoami 或 dir)。", + ShortDescription: "测试 WebShell 连接", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "connection_id": map[string]interface{}{ + "type": "string", + "description": "要测试的 WebShell 连接 ID(必填)", + }, + "command": map[string]interface{}{ + "type": "string", + "description": "测试命令,默认为 whoami(Linux)或 dir(Windows)", + }, + }, + "required": []string{"connection_id"}, + }, + } + testHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + connID, _ := args["connection_id"].(string) + if connID == "" { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "错误: connection_id 参数必填"}}, + IsError: true, + }, nil + } + + // 获取连接 + conn, err := db.GetWebshellConnection(connID) + if err != nil || conn == nil { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: "未找到指定的 WebShell 连接: " + connID}}, + IsError: true, + }, nil + } + + // 确定测试命令 + testCmd, _ := args["command"].(string) + if testCmd == "" { + // 根据 shell 类型选择默认命令 + if conn.Type == "asp" || conn.Type == "aspx" { + testCmd = "dir" + } else { + testCmd = "whoami" + } + } + + // 执行测试命令 + output, ok, errMsg := webshellHandler.ExecWithConnection(conn, testCmd) + if errMsg != "" { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: fmt.Sprintf("连接测试失败!\n\n连接ID: %s\nURL: %s\n错误: %s", connID, conn.URL, errMsg)}}, + IsError: true, + }, nil + } + + if !ok { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: fmt.Sprintf("连接测试失败!HTTP 非 200\n\n连接ID: %s\nURL: %s\n输出: %s", connID, conn.URL, output)}}, + IsError: true, + }, nil + } + + return &mcp.ToolResult{ + Content: []mcp.Content{{ + Type: "text", + Text: fmt.Sprintf("连接测试成功!\n\n连接ID: %s\nURL: %s\n类型: %s\n\n测试命令: %s\n输出结果:\n%s", connID, conn.URL, conn.Type, testCmd, output), + }}, + IsError: false, + }, nil + } + mcpServer.RegisterTool(testTool, testHandler) + + logger.Debug("WebShell 管理工具注册成功") +} + +// initializeKnowledge 初始化知识库组件(用于动态初始化) +func initializeKnowledge( + cfg *config.Config, + db *database.DB, + knowledgeDBConn *database.DB, + mcpServer *mcp.Server, + agentHandler *handler.AgentHandler, + app *App, // 传递 App 引用以便更新知识库组件 + logger *zap.Logger, +) (*handler.KnowledgeHandler, error) { + // 确定知识库数据库路径 + knowledgeDBPath := cfg.Database.KnowledgeDBPath + var knowledgeDB *sql.DB + + if knowledgeDBPath != "" { + // 使用独立的知识库数据库 + // 确保目录存在 + if err := os.MkdirAll(filepath.Dir(knowledgeDBPath), 0755); err != nil { + return nil, fmt.Errorf("创建知识库数据库目录失败: %w", err) + } + + var err error + knowledgeDBConn, err = database.NewKnowledgeDB(knowledgeDBPath, logger) + if err != nil { + return nil, fmt.Errorf("初始化知识库数据库失败: %w", err) + } + knowledgeDB = knowledgeDBConn.DB + logger.Info("使用独立的知识库数据库", zap.String("path", knowledgeDBPath)) + } else { + // 向后兼容:使用会话数据库 + knowledgeDB = db.DB + logger.Info("使用会话数据库存储知识库数据(建议配置knowledge_db_path以分离数据)") + } + + // 创建知识库管理器 + knowledgeManager := knowledge.NewManager(knowledgeDB, cfg.Knowledge.BasePath, logger) + + // 创建嵌入器 + // 使用OpenAI配置的API Key(如果知识库配置中没有指定) + if cfg.Knowledge.Embedding.APIKey == "" { + cfg.Knowledge.Embedding.APIKey = cfg.OpenAI.APIKey + } + if cfg.Knowledge.Embedding.BaseURL == "" { + cfg.Knowledge.Embedding.BaseURL = cfg.OpenAI.BaseURL + } + + embedder, err := knowledge.NewEmbedder(context.Background(), &cfg.Knowledge, &cfg.OpenAI, logger) + if err != nil { + return nil, fmt.Errorf("初始化知识库嵌入器失败: %w", err) + } + + // 创建检索器(Eino MultiQuery + 重排流水线) + retrievalConfig := knowledge.RetrievalConfigFromYAML(cfg.Knowledge.Retrieval) + knowledgeRetriever := knowledge.NewRetriever(knowledgeDB, embedder, retrievalConfig, logger) + if err := knowledge.WireRetrieverPipeline(context.Background(), knowledgeRetriever, &cfg.OpenAI); err != nil { + return nil, fmt.Errorf("初始化知识库检索流水线失败: %w", err) + } + + // 创建索引器(Eino Compose 链) + knowledgeIndexer, err := knowledge.NewIndexer(context.Background(), knowledgeDB, embedder, logger, &cfg.Knowledge) + if err != nil { + return nil, fmt.Errorf("初始化知识库索引器失败: %w", err) + } + + // 注册知识检索工具到MCP服务器 + knowledge.RegisterKnowledgeTool(mcpServer, knowledgeRetriever, knowledgeManager, logger) + + // 创建知识库API处理器 + knowledgeHandler := handler.NewKnowledgeHandler(knowledgeManager, knowledgeRetriever, knowledgeIndexer, db, logger) + if app != nil && app.auditSvc != nil { + knowledgeHandler.SetAudit(app.auditSvc) + } + logger.Info("知识库模块初始化完成", zap.Bool("handler_created", knowledgeHandler != nil)) + + // 设置知识库管理器到AgentHandler以便记录检索日志 + agentHandler.SetKnowledgeManager(knowledgeManager) + + // 更新 App 中的知识库组件(如果 App 不为 nil,说明是动态初始化) + if app != nil { + app.knowledgeManager = knowledgeManager + app.knowledgeRetriever = knowledgeRetriever + app.knowledgeIndexer = knowledgeIndexer + app.knowledgeHandler = knowledgeHandler + // 如果使用独立数据库,更新 knowledgeDB + if knowledgeDBPath != "" { + app.knowledgeDB = knowledgeDBConn + } + logger.Info("App 中的知识库组件已更新") + } + + // 扫描知识库并建立索引(异步) + go func() { + itemsToIndex, err := knowledgeManager.ScanKnowledgeBase() + if err != nil { + logger.Warn("扫描知识库失败", zap.Error(err)) + return + } + + // 检查是否已有索引 + hasIndex, err := knowledgeIndexer.HasIndex() + if err != nil { + logger.Warn("检查索引状态失败", zap.Error(err)) + return + } + + if hasIndex { + // 如果已有索引,只索引新添加或更新的项 + if len(itemsToIndex) > 0 { + logger.Info("检测到已有知识库索引,开始增量索引", zap.Int("count", len(itemsToIndex))) + ctx := context.Background() + consecutiveFailures := 0 + var firstFailureItemID string + var firstFailureError error + failedCount := 0 + + for _, itemID := range itemsToIndex { + if err := knowledgeIndexer.IndexItem(ctx, itemID); err != nil { + failedCount++ + consecutiveFailures++ + + if consecutiveFailures == 1 { + firstFailureItemID = itemID + firstFailureError = err + logger.Warn("索引知识项失败", zap.String("itemId", itemID), zap.Error(err)) + } + + // 如果连续失败2次,立即停止增量索引 + if consecutiveFailures >= 2 { + logger.Error("连续索引失败次数过多,立即停止增量索引", + zap.Int("consecutiveFailures", consecutiveFailures), + zap.Int("totalItems", len(itemsToIndex)), + zap.String("firstFailureItemId", firstFailureItemID), + zap.Error(firstFailureError), + ) + break + } + continue + } + + // 成功时重置连续失败计数 + if consecutiveFailures > 0 { + consecutiveFailures = 0 + firstFailureItemID = "" + firstFailureError = nil + } + } + logger.Info("增量索引完成", zap.Int("totalItems", len(itemsToIndex)), zap.Int("failedCount", failedCount)) + } else { + logger.Info("检测到已有知识库索引,没有需要索引的新项或更新项") + } + return + } + + // 冷启动:仅为尚无向量的知识项构建索引(与 IndexMissing 语义一致) + logger.Info("未检测到知识库索引,开始自动构建索引") + ctx := context.Background() + if err := knowledgeIndexer.IndexMissing(ctx); err != nil { + logger.Warn("自动构建知识库索引失败", zap.Error(err)) + } + }() + + return knowledgeHandler, nil +} + +// corsMiddleware allows same-origin requests, valid Chromium extension +// origins, and exact origins explicitly configured by the operator. CORS is +// not an authentication boundary; API access still requires a valid session. +func corsMiddleware(configuredOrigins []string) gin.HandlerFunc { + allowedOrigins := make(map[string]struct{}, len(configuredOrigins)) + for _, origin := range configuredOrigins { + if normalized, ok := normalizeCORSOrigin(origin); ok { + allowedOrigins[normalized] = struct{}{} + } + } + + return func(c *gin.Context) { + origin := strings.TrimSpace(c.GetHeader("Origin")) + if origin != "" { + c.Writer.Header().Add("Vary", "Origin") + normalized, valid := normalizeCORSOrigin(origin) + _, explicitlyAllowed := allowedOrigins[normalized] + parsed, _ := url.Parse(origin) + sameHost := valid && strings.EqualFold(parsed.Host, c.Request.Host) + browserExtension := valid && isChromiumExtensionOrigin(parsed) + if !sameHost && !browserExtension && !explicitlyAllowed { + c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error": "cross-origin request denied"}) + return + } + c.Writer.Header().Set("Access-Control-Allow-Origin", origin) + c.Writer.Header().Set("Access-Control-Allow-Credentials", "true") + } + c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization, accept, origin, Cache-Control, X-Requested-With") + c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE") + c.Writer.Header().Set("Access-Control-Max-Age", "600") + + if c.Request.Method == "OPTIONS" { + c.AbortWithStatus(204) + return + } + + c.Next() + } +} + +// isChromiumExtensionOrigin accepts only Chrome's canonical 32-character +// extension IDs (letters a-p). It does not allow arbitrary custom schemes or +// web origins, and the extension must separately obtain host permission. +func isChromiumExtensionOrigin(origin *url.URL) bool { + if origin == nil || !strings.EqualFold(origin.Scheme, "chrome-extension") || origin.Port() != "" { + return false + } + id := strings.ToLower(origin.Hostname()) + if len(id) != 32 { + return false + } + for _, ch := range id { + if ch < 'a' || ch > 'p' { + return false + } + } + return true +} + +// normalizeCORSOrigin validates and canonicalizes a serialized origin. CORS +// origins never contain credentials, paths, query strings, or fragments. +func normalizeCORSOrigin(raw string) (string, bool) { + raw = strings.TrimSpace(raw) + if raw == "" || raw == "*" || strings.EqualFold(raw, "null") { + return "", false + } + parsed, err := url.Parse(raw) + if err != nil || parsed.Scheme == "" || parsed.Host == "" || parsed.User != nil || + (parsed.Path != "" && parsed.Path != "/") || parsed.RawQuery != "" || parsed.Fragment != "" { + return "", false + } + return strings.ToLower(parsed.Scheme) + "://" + strings.ToLower(parsed.Host), true +} diff --git a/internal/app/asset_tools.go b/internal/app/asset_tools.go new file mode 100644 index 00000000..611976d7 --- /dev/null +++ b/internal/app/asset_tools.go @@ -0,0 +1,511 @@ +package app + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "strings" + "time" + + "cyberstrike-ai/internal/authctx" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" + + "go.uber.org/zap" +) + +const agentAssetPageSizeMax = 50 + +func registerAssetTools(server *mcp.Server, db *database.DB, logger *zap.Logger) { + if server == nil || db == nil { + return + } + properties := assetMutationProperties() + + server.RegisterTool(mcp.Tool{ + Name: builtin.ToolCreateAsset, ShortDescription: "新增或去重更新资产", + Description: "向资产库新增资产。按目标+端口+协议去重;若资产已存在则更新非空字段。至少提供 host、ip、domain 之一。", + // Bedrock rejects tool schemas with top-level oneOf/allOf/anyOf. The + // host/ip/domain requirement is enforced by assetFromCreateArgs below. + InputSchema: map[string]interface{}{"type": "object", "properties": properties}, + }, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + asset, err := assetFromCreateArgs(args) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + access, owner, global := assetAccessFromToolContext(ctx, "asset:write") + result, err := db.UpsertAssets([]*database.Asset{asset}, owner, global) + if err != nil { + logger.Error("Agent 保存资产失败", zap.Error(err)) + return textResult("错误: "+err.Error(), true), nil + } + if result.Skipped > 0 || asset.ID == "" { + return textResult("资产未保存:同一资产已存在但当前用户无权更新,或目标字段为空", true), nil + } + saved, err := db.GetAsset(asset.ID, access) + if err != nil { + return textResult("资产已保存,但无法读取结果: "+err.Error(), true), nil + } + action := "created" + if result.Updated > 0 { + action = "updated" + } + return assetJSONResult(map[string]interface{}{"action": action, "asset": assetToolDetail(saved)}) + }) + + server.RegisterTool(mcp.Tool{ + Name: builtin.ToolGetAsset, ShortDescription: "按 ID 查看资产详情", Description: "按资产 ID 返回完整资产详情。查询列表时先用 query_assets,避免一次拉取过多详情。", + InputSchema: map[string]interface{}{"type": "object", "properties": map[string]interface{}{"id": map[string]interface{}{"type": "string", "description": "资产 ID"}}, "required": []string{"id"}}, + }, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + projectID, projectScoped, err := agentAssetProjectScope(db, ctx) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + asset, err := db.GetAsset(strings.TrimSpace(strArg(args, "id")), assetAccessOnly(ctx, "asset:read")) + if err != nil { + if err == sql.ErrNoRows { + return textResult("错误: 资产不存在或无权查看", true), nil + } + return textResult("错误: "+err.Error(), true), nil + } + if projectScoped && strings.TrimSpace(asset.ProjectID) != projectID { + return textResult("错误: 资产不存在或不属于当前对话绑定的项目", true), nil + } + return assetJSONResult(assetToolDetail(asset)) + }) + + server.RegisterTool(mcp.Tool{ + Name: builtin.ToolQueryAssets, ShortDescription: "灵活分页查询资产", + Description: "分页查询资产。支持精确字段、时间范围、扫描状态和白名单排序。查最久未扫描资产请使用 sort_by=last_scan_at、sort_order=asc;从未扫描资产会排在最前。默认每页 20 条,最大 50 条,返回精简摘要;使用 get_asset 获取单条详情。", + InputSchema: assetQuerySchema(), + }, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + filter, page, pageSize, err := assetFilterFromToolArgs(args) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + projectID, projectScoped, err := agentAssetProjectScope(db, ctx) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + if projectScoped { + // 对话绑定项目后,项目范围是服务端强制边界;不能通过工具参数扩大或切换范围。 + filter.ProjectID = projectID + } + items, total, err := db.ListAssets(pageSize, (page-1)*pageSize, filter, assetAccessOnly(ctx, "asset:read")) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + totalPages := (total + pageSize - 1) / pageSize + if totalPages < 1 { + totalPages = 1 + } + var b strings.Builder + b.WriteString(fmt.Sprintf("资产查询:第 %d/%d 页,本页 %d 条,共 %d 条,page_size=%d\n", page, totalPages, len(items), total, pageSize)) + for _, asset := range items { + b.WriteString(formatAssetListItem(asset)) + b.WriteByte('\n') + } + if page < totalPages { + b.WriteString(fmt.Sprintf("下一页:保持筛选条件并设置 page=%d。", page+1)) + } + return textResult(b.String(), false), nil + }) + + updateProperties := assetMutationProperties() + updateProperties["id"] = map[string]interface{}{"type": "string", "description": "资产 ID"} + server.RegisterTool(mcp.Tool{ + Name: builtin.ToolUpdateAsset, ShortDescription: "局部更新资产", + Description: "按 ID 局部更新资产,只修改显式传入的字段;可传空 project_id 清除项目绑定,可传空 tags 清空标签。", + InputSchema: map[string]interface{}{"type": "object", "properties": updateProperties, "required": []string{"id"}}, + }, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + id := strings.TrimSpace(strArg(args, "id")) + access := assetAccessOnly(ctx, "asset:write") + asset, err := db.GetAsset(id, access) + if err != nil { + return textResult("错误: 资产不存在或无权更新", true), nil + } + if err := applyAssetPatch(asset, args); err != nil { + return textResult("错误: "+err.Error(), true), nil + } + if err := db.UpdateAsset(id, asset, access); err != nil { + return textResult("错误: "+err.Error(), true), nil + } + updated, err := db.GetAsset(id, access) + if err != nil { + return textResult("资产已更新,但无法读取结果: "+err.Error(), true), nil + } + return assetJSONResult(map[string]interface{}{"action": "updated", "asset": assetToolDetail(updated)}) + }) + + server.RegisterTool(mcp.Tool{ + Name: builtin.ToolDeleteAsset, ShortDescription: "删除资产", Description: "按 ID 永久删除资产记录。仅在用户明确要求删除时调用。", + InputSchema: map[string]interface{}{"type": "object", "properties": map[string]interface{}{"id": map[string]interface{}{"type": "string", "description": "资产 ID"}}, "required": []string{"id"}}, + }, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + id := strings.TrimSpace(strArg(args, "id")) + if err := db.DeleteAsset(id, assetAccessOnly(ctx, "asset:delete")); err != nil { + return textResult("错误: 资产不存在或无权删除", true), nil + } + return textResult("资产已删除: "+id, false), nil + }) + + server.RegisterTool(mcp.Tool{ + Name: builtin.ToolCompleteAssetScan, + ShortDescription: "完成资产扫描并回写结果", + Description: "目标扫描完成后调用:把资产的上次扫描时间更新为当前时间,并关联当前对话。相关漏洞数量不手填,而是自动统计当前扫描对话中通过 record_vulnerability 保存的漏洞。应在漏洞均已落库后调用;一个扫描对话建议只对应一个资产。", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "id": map[string]interface{}{"type": "string", "description": "已完成扫描的资产 ID"}, + }, + "required": []string{"id"}, + }, + }, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + id := strings.TrimSpace(strArg(args, "id")) + conversationID := conversationIDFromToolCtx(ctx) + if conversationID == "" { + return textResult("错误: 无法确定当前扫描对话", true), nil + } + access := assetAccessOnly(ctx, "asset:write") + if err := db.CompleteAssetScan(id, conversationID, access); err != nil { + if err == sql.ErrNoRows { + return textResult("错误: 资产不存在或无权回写扫描结果", true), nil + } + return textResult("错误: "+err.Error(), true), nil + } + updated, err := db.GetAsset(id, access) + if err != nil { + return textResult("扫描结果已回写,但无法读取资产: "+err.Error(), true), nil + } + return assetJSONResult(map[string]interface{}{ + "action": "scan_completed", + "message": "上次扫描时间已更新;相关漏洞数由当前扫描对话中已保存的漏洞自动计算", + "asset": assetToolDetail(updated), + }) + }) +} + +func assetMutationProperties() map[string]interface{} { + return map[string]interface{}{ + "project_id": map[string]interface{}{"type": "string"}, "host": map[string]interface{}{"type": "string"}, + "ip": map[string]interface{}{"type": "string"}, "port": map[string]interface{}{"type": "integer", "minimum": 0, "maximum": 65535}, + "domain": map[string]interface{}{"type": "string"}, "protocol": map[string]interface{}{"type": "string"}, + "title": map[string]interface{}{"type": "string"}, "server": map[string]interface{}{"type": "string"}, + "country": map[string]interface{}{"type": "string"}, "province": map[string]interface{}{"type": "string"}, "city": map[string]interface{}{"type": "string"}, + "responsible_person": map[string]interface{}{"type": "string"}, "department": map[string]interface{}{"type": "string"}, + "business_system": map[string]interface{}{"type": "string"}, + "environment": map[string]interface{}{"type": "string", "enum": []string{"production", "staging", "testing", "development", "other"}}, + "criticality": map[string]interface{}{"type": "string", "enum": []string{"critical", "high", "medium", "low"}}, + "source": map[string]interface{}{"type": "string"}, "source_query": map[string]interface{}{"type": "string"}, + "status": map[string]interface{}{"type": "string", "enum": []string{"active", "inactive"}}, + "tags": map[string]interface{}{"type": "array", "items": map[string]interface{}{"type": "string"}, "maxItems": 50}, + } +} + +func assetQuerySchema() map[string]interface{} { + properties := map[string]interface{}{ + "q": map[string]interface{}{"type": "string", "description": "模糊搜索 host、IP、域名、标题、服务和标签"}, + "project_id": map[string]interface{}{"type": "string"}, "status": map[string]interface{}{"type": "string", "enum": []string{"active", "inactive"}}, + "protocol": map[string]interface{}{"type": "string"}, "source": map[string]interface{}{"type": "string"}, "tag": map[string]interface{}{"type": "string"}, + "host": map[string]interface{}{"type": "string"}, "ip": map[string]interface{}{"type": "string"}, "domain": map[string]interface{}{"type": "string"}, + "port": map[string]interface{}{"type": "integer", "minimum": 0, "maximum": 65535}, + "risk_level": map[string]interface{}{"type": "string", "enum": []string{"unassessed", "critical", "high", "medium", "low", "info", "normal"}}, + "min_vulnerabilities": map[string]interface{}{"type": "integer", "minimum": 0}, + "max_vulnerabilities": map[string]interface{}{"type": "integer", "minimum": 0}, + "country": map[string]interface{}{"type": "string"}, "province": map[string]interface{}{"type": "string"}, "city": map[string]interface{}{"type": "string"}, + "responsible_person": map[string]interface{}{"type": "string"}, "department": map[string]interface{}{"type": "string"}, + "business_system": map[string]interface{}{"type": "string"}, "environment": map[string]interface{}{"type": "string"}, "criticality": map[string]interface{}{"type": "string"}, + "scan_state": map[string]interface{}{"type": "string", "enum": []string{"never", "scanned"}, "description": "never=从未扫描,scanned=扫描过"}, + "scan_overdue_days": map[string]interface{}{"type": "integer", "minimum": 1}, + "last_scan_before": map[string]interface{}{"type": "string", "description": "RFC3339 时间或 YYYY-MM-DD"}, + "last_scan_after": map[string]interface{}{"type": "string", "description": "RFC3339 时间或 YYYY-MM-DD"}, + "first_seen_before": map[string]interface{}{"type": "string", "description": "RFC3339 时间或 YYYY-MM-DD"}, + "first_seen_after": map[string]interface{}{"type": "string", "description": "RFC3339 时间或 YYYY-MM-DD"}, + "last_seen_before": map[string]interface{}{"type": "string", "description": "RFC3339 时间或 YYYY-MM-DD"}, + "last_seen_after": map[string]interface{}{"type": "string", "description": "RFC3339 时间或 YYYY-MM-DD"}, + "sort_by": map[string]interface{}{"type": "string", "enum": []string{"last_seen_at", "last_scan_at", "first_seen_at", "created_at", "updated_at", "host", "port", "risk_level", "vulnerability_count"}}, + "sort_order": map[string]interface{}{"type": "string", "enum": []string{"asc", "desc"}}, + "page": map[string]interface{}{"type": "integer", "minimum": 1}, + "page_size": map[string]interface{}{"type": "integer", "minimum": 1, "maximum": agentAssetPageSizeMax}, + } + return map[string]interface{}{"type": "object", "properties": properties} +} + +func assetFromCreateArgs(args map[string]interface{}) (*database.Asset, error) { + asset := &database.Asset{} + if err := applyAssetPatch(asset, args); err != nil { + return nil, err + } + if strings.TrimSpace(asset.Host) == "" && strings.TrimSpace(asset.IP) == "" && strings.TrimSpace(asset.Domain) == "" { + return nil, fmt.Errorf("host、ip、domain 至少需要一个") + } + return asset, nil +} + +func applyAssetPatch(asset *database.Asset, args map[string]interface{}) error { + setString := func(key string, dst *string) { + if _, ok := args[key]; ok { + *dst = strings.TrimSpace(strArg(args, key)) + } + } + setString("project_id", &asset.ProjectID) + setString("host", &asset.Host) + setString("ip", &asset.IP) + setString("domain", &asset.Domain) + setString("protocol", &asset.Protocol) + setString("title", &asset.Title) + setString("server", &asset.Server) + setString("country", &asset.Country) + setString("province", &asset.Province) + setString("city", &asset.City) + setString("responsible_person", &asset.ResponsiblePerson) + setString("department", &asset.Department) + setString("business_system", &asset.BusinessSystem) + setString("environment", &asset.Environment) + setString("criticality", &asset.Criticality) + setString("source", &asset.Source) + setString("source_query", &asset.SourceQuery) + setString("status", &asset.Status) + if _, ok := args["port"]; ok { + port := intArg(args, "port", -1) + if port < 0 || port > 65535 { + return fmt.Errorf("port 必须在 0-65535 之间") + } + asset.Port = port + } + if raw, ok := args["tags"]; ok { + tags, err := stringSliceArg(raw) + if err != nil { + return fmt.Errorf("tags: %w", err) + } + asset.Tags = tags + } + return nil +} + +func assetFilterFromToolArgs(args map[string]interface{}) (database.AssetListFilter, int, int, error) { + filter := database.AssetListFilter{ + Search: strings.TrimSpace(strArg(args, "q")), ProjectID: strings.TrimSpace(strArg(args, "project_id")), Status: strings.ToLower(strings.TrimSpace(strArg(args, "status"))), + Protocol: strings.ToLower(strings.TrimSpace(strArg(args, "protocol"))), Source: strings.TrimSpace(strArg(args, "source")), Tag: strings.TrimSpace(strArg(args, "tag")), + Host: strings.TrimSpace(strArg(args, "host")), IP: strings.TrimSpace(strArg(args, "ip")), Domain: strings.TrimSpace(strArg(args, "domain")), + ScanState: strings.ToLower(strings.TrimSpace(strArg(args, "scan_state"))), SortBy: strings.ToLower(strings.TrimSpace(strArg(args, "sort_by"))), + SortOrder: strings.ToLower(strings.TrimSpace(strArg(args, "sort_order"))), + RiskLevel: strings.ToLower(strings.TrimSpace(strArg(args, "risk_level"))), + Country: strings.TrimSpace(strArg(args, "country")), Province: strings.TrimSpace(strArg(args, "province")), City: strings.TrimSpace(strArg(args, "city")), + ResponsiblePerson: strings.TrimSpace(strArg(args, "responsible_person")), Department: strings.TrimSpace(strArg(args, "department")), + BusinessSystem: strings.TrimSpace(strArg(args, "business_system")), Environment: strings.ToLower(strings.TrimSpace(strArg(args, "environment"))), + Criticality: strings.ToLower(strings.TrimSpace(strArg(args, "criticality"))), + } + if !oneOfOrEmpty(filter.Status, "active", "inactive") { + return filter, 0, 0, fmt.Errorf("status 仅支持 active 或 inactive") + } + if !oneOfOrEmpty(filter.ScanState, "never", "scanned") { + return filter, 0, 0, fmt.Errorf("scan_state 仅支持 never 或 scanned") + } + if !oneOfOrEmpty(filter.SortBy, "last_seen_at", "last_scan_at", "first_seen_at", "created_at", "updated_at", "host", "port", "risk_level", "vulnerability_count") { + return filter, 0, 0, fmt.Errorf("sort_by 不受支持") + } + if !oneOfOrEmpty(filter.SortOrder, "asc", "desc") { + return filter, 0, 0, fmt.Errorf("sort_order 仅支持 asc 或 desc") + } + if _, ok := args["port"]; ok { + port := intArg(args, "port", -1) + if port < 0 || port > 65535 { + return filter, 0, 0, fmt.Errorf("port 必须在 0-65535 之间") + } + filter.Port = &port + } + if _, ok := args["min_vulnerabilities"]; ok { + value := intArg(args, "min_vulnerabilities", -1) + if value < 0 { + return filter, 0, 0, fmt.Errorf("min_vulnerabilities 不能小于 0") + } + filter.MinVulnerabilities = &value + } + if _, ok := args["max_vulnerabilities"]; ok { + value := intArg(args, "max_vulnerabilities", -1) + if value < 0 { + return filter, 0, 0, fmt.Errorf("max_vulnerabilities 不能小于 0") + } + filter.MaxVulnerabilities = &value + } + if _, ok := args["scan_overdue_days"]; ok { + value := intArg(args, "scan_overdue_days", 0) + if value < 1 { + return filter, 0, 0, fmt.Errorf("scan_overdue_days 必须大于 0") + } + filter.ScanOverdueDays = &value + } + var err error + if filter.LastScanBefore, err = parseAssetToolTime("last_scan_before", strArg(args, "last_scan_before")); err != nil { + return filter, 0, 0, err + } + if filter.LastScanAfter, err = parseAssetToolTime("last_scan_after", strArg(args, "last_scan_after")); err != nil { + return filter, 0, 0, err + } + if filter.FirstSeenBefore, err = parseAssetToolTime("first_seen_before", strArg(args, "first_seen_before")); err != nil { + return filter, 0, 0, err + } + if filter.FirstSeenAfter, err = parseAssetToolTime("first_seen_after", strArg(args, "first_seen_after")); err != nil { + return filter, 0, 0, err + } + if filter.LastSeenBefore, err = parseAssetToolTime("last_seen_before", strArg(args, "last_seen_before")); err != nil { + return filter, 0, 0, err + } + if filter.LastSeenAfter, err = parseAssetToolTime("last_seen_after", strArg(args, "last_seen_after")); err != nil { + return filter, 0, 0, err + } + page := intArg(args, "page", 1) + pageSize := intArg(args, "page_size", 20) + if page < 1 || page > 1_000_000 { + return filter, 0, 0, fmt.Errorf("page 必须在 1-1000000 之间") + } + if pageSize < 1 || pageSize > agentAssetPageSizeMax { + return filter, 0, 0, fmt.Errorf("page_size 必须在 1-%d 之间", agentAssetPageSizeMax) + } + return filter, page, pageSize, nil +} + +func oneOfOrEmpty(value string, allowed ...string) bool { + if value == "" { + return true + } + for _, candidate := range allowed { + if value == candidate { + return true + } + } + return false +} + +func parseAssetToolTime(field, value string) (*time.Time, error) { + value = strings.TrimSpace(value) + if value == "" { + return nil, nil + } + for _, layout := range []string{time.RFC3339, "2006-01-02"} { + if parsed, err := time.Parse(layout, value); err == nil { + return &parsed, nil + } + } + return nil, fmt.Errorf("%s 必须是 RFC3339 时间或 YYYY-MM-DD", field) +} + +func stringSliceArg(raw interface{}) ([]string, error) { + values := []string{} + switch typed := raw.(type) { + case []string: + values = append(values, typed...) + case []interface{}: + for _, item := range typed { + value, ok := item.(string) + if !ok { + return nil, fmt.Errorf("必须是字符串数组") + } + values = append(values, value) + } + default: + return nil, fmt.Errorf("必须是字符串数组") + } + if len(values) > 50 { + return nil, fmt.Errorf("最多 50 个标签") + } + return values, nil +} + +func assetAccessOnly(ctx context.Context, permission string) database.RBACListAccess { + principal, ok := authctx.PrincipalFromContext(ctx) + if !ok { + return database.RBACListAccess{} + } + return database.RBACListAccess{UserID: principal.UserID, Scope: principal.ScopeFor(permission)} +} + +func assetAccessFromToolContext(ctx context.Context, permission string) (database.RBACListAccess, string, bool) { + principal, ok := authctx.PrincipalFromContext(ctx) + if !ok { + return database.RBACListAccess{}, "", false + } + access := database.RBACListAccess{UserID: principal.UserID, Scope: principal.ScopeFor(permission)} + return access, principal.UserID, access.Scope == database.RBACScopeAll +} + +// agentAssetProjectScope returns the hard asset-read boundary implied by the +// current conversation. An unbound conversation (or a tool call outside a +// conversation) keeps the existing all-accessible-assets behavior. A bound +// conversation can only read assets assigned to that exact project. +func agentAssetProjectScope(db *database.DB, ctx context.Context) (projectID string, scoped bool, err error) { + conversationID := conversationIDFromToolCtx(ctx) + if conversationID == "" { + return "", false, nil + } + projectID, err = db.GetConversationProjectID(conversationID) + if err != nil { + return "", false, fmt.Errorf("无法确定当前对话的项目范围") + } + projectID = strings.TrimSpace(projectID) + return projectID, projectID != "", nil +} + +func formatAssetListItem(asset *database.Asset) string { + target := asset.Domain + if target == "" { + target = asset.IP + } + if target == "" { + target = asset.Host + } + if asset.Port > 0 { + target = fmt.Sprintf("%s:%d", target, asset.Port) + } + lastScan := "never" + if asset.LastScanAt != nil { + lastScan = asset.LastScanAt.Format(time.RFC3339) + } + return fmt.Sprintf("- id=%s | target=%s | protocol=%s | status=%s | last_scan_at=%s | risk=%s | vulnerabilities=%d", asset.ID, truncateRunes(target, 120), truncateRunes(asset.Protocol, 30), truncateRunes(asset.Status, 30), lastScan, asset.RiskLevel, asset.VulnerabilityCount) +} + +// assetToolDetail keeps even a single unusually large imported record from +// consuming the model context. The database and HTTP API retain full values. +func assetToolDetail(asset *database.Asset) map[string]interface{} { + if asset == nil { + return nil + } + tags := make([]string, 0, len(asset.Tags)) + for i, tag := range asset.Tags { + if i >= 50 { + break + } + tags = append(tags, truncateRunes(tag, 100)) + } + detail := map[string]interface{}{ + "id": asset.ID, "project_id": asset.ProjectID, "project_name": truncateRunes(asset.ProjectName, 200), + "host": truncateRunes(asset.Host, 500), "ip": truncateRunes(asset.IP, 100), "port": asset.Port, + "domain": truncateRunes(asset.Domain, 255), "protocol": truncateRunes(asset.Protocol, 50), + "title": truncateRunes(asset.Title, 500), "server": truncateRunes(asset.Server, 500), + "country": truncateRunes(asset.Country, 100), "province": truncateRunes(asset.Province, 100), "city": truncateRunes(asset.City, 100), + "responsible_person": truncateRunes(asset.ResponsiblePerson, 255), "department": truncateRunes(asset.Department, 255), + "business_system": truncateRunes(asset.BusinessSystem, 255), "environment": asset.Environment, "criticality": asset.Criticality, + "source": truncateRunes(asset.Source, 100), "source_query": truncateRunes(asset.SourceQuery, 2000), + "status": truncateRunes(asset.Status, 50), "tags": tags, + "first_seen_at": asset.FirstSeenAt, "last_seen_at": asset.LastSeenAt, "created_at": asset.CreatedAt, "updated_at": asset.UpdatedAt, + "last_scan_conversation_id": asset.LastScanConversationID, "last_scan_queue_id": asset.LastScanQueueID, "last_scan_task_id": asset.LastScanTaskID, + "vulnerability_count": asset.VulnerabilityCount, "risk_level": asset.RiskLevel, + } + if asset.LastScanAt != nil { + detail["last_scan_at"] = asset.LastScanAt + } + if len(asset.Tags) > len(tags) { + detail["tags_truncated"] = true + } + return detail +} + +func assetJSONResult(value interface{}) (*mcp.ToolResult, error) { + encoded, err := json.MarshalIndent(value, "", " ") + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + return textResult(string(encoded), false), nil +} diff --git a/internal/app/asset_tools_test.go b/internal/app/asset_tools_test.go new file mode 100644 index 00000000..9cb491b7 --- /dev/null +++ b/internal/app/asset_tools_test.go @@ -0,0 +1,201 @@ +package app + +import ( + "context" + "path/filepath" + "strings" + "testing" + + "cyberstrike-ai/internal/authctx" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" + + "go.uber.org/zap" +) + +func TestAssetToolsCRUDQueryAndPageLimit(t *testing.T) { + db, err := database.NewDB(filepath.Join(t.TempDir(), "asset-tools.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + user, err := db.CreateRBACUser("asset-agent", "Asset Agent", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + principal := authctx.NewPrincipal(user.ID, user.Username, database.RBACScopeAssigned, map[string]bool{ + "asset:read": true, "asset:write": true, "asset:delete": true, + }) + ctx := authctx.WithPrincipal(context.Background(), principal) + server := mcp.NewServer(zap.NewNop()) + server.SetToolAuthorizer(mcpToolAuthorizer(db)) + registerAssetTools(server, db, zap.NewNop()) + + wantTools := map[string]bool{ + builtin.ToolCreateAsset: false, builtin.ToolGetAsset: false, builtin.ToolQueryAssets: false, + builtin.ToolUpdateAsset: false, builtin.ToolDeleteAsset: false, builtin.ToolCompleteAssetScan: false, + } + for _, tool := range server.GetAllTools() { + if _, ok := wantTools[tool.Name]; ok { + wantTools[tool.Name] = true + } + } + for name, found := range wantTools { + if !found { + t.Fatalf("asset tool not registered: %s", name) + } + } + + for _, tool := range server.GetAllTools() { + if tool.Name != builtin.ToolCreateAsset { + continue + } + for _, keyword := range []string{"oneOf", "allOf", "anyOf"} { + if _, exists := tool.InputSchema[keyword]; exists { + t.Fatalf("create asset schema contains Bedrock-incompatible top-level %s", keyword) + } + } + } + + result, _, err := server.CallTool(ctx, builtin.ToolCreateAsset, map[string]interface{}{"title": "Missing target"}) + if err != nil || result == nil || !result.IsError { + t.Fatalf("create asset accepted missing host/ip/domain: result=%#v err=%v", result, err) + } + + result, _, err = server.CallTool(ctx, builtin.ToolCreateAsset, map[string]interface{}{ + "ip": "192.0.2.42", "port": 443, "protocol": "https", "title": "Before", "tags": []interface{}{"prod"}, + }) + if err != nil || result == nil || result.IsError { + t.Fatalf("create asset result=%#v err=%v", result, err) + } + assets, total, err := db.ListAssets(20, 0, database.AssetListFilter{}, database.RBACListAccess{UserID: user.ID, Scope: database.RBACScopeAssigned}) + if err != nil || total != 1 || len(assets) != 1 { + t.Fatalf("saved assets total=%d len=%d err=%v", total, len(assets), err) + } + id := assets[0].ID + + result, _, err = server.CallTool(ctx, builtin.ToolUpdateAsset, map[string]interface{}{"id": id, "title": "After"}) + if err != nil || result == nil || result.IsError { + t.Fatalf("update asset result=%#v err=%v", result, err) + } + updated, err := db.GetAsset(id, database.RBACListAccess{UserID: user.ID, Scope: database.RBACScopeAssigned}) + if err != nil || updated.Title != "After" || updated.IP != "192.0.2.42" { + t.Fatalf("partial update lost fields: %#v err=%v", updated, err) + } + + result, _, err = server.CallTool(ctx, builtin.ToolQueryAssets, map[string]interface{}{ + "sort_by": "last_scan_at", "sort_order": "asc", "page": 1, "page_size": 1, + }) + if err != nil || result == nil || result.IsError || !strings.Contains(toolResultText(result), "第 1/1 页") || !strings.Contains(toolResultText(result), "last_scan_at=never") { + t.Fatalf("query asset result=%#v err=%v", result, err) + } + result, _, err = server.CallTool(ctx, builtin.ToolQueryAssets, map[string]interface{}{"page_size": agentAssetPageSizeMax + 1}) + if err != nil || result == nil || !result.IsError { + t.Fatalf("oversized page was accepted: result=%#v err=%v", result, err) + } + + conversation, err := db.CreateConversation("asset scan", database.ConversationCreateMeta{}) + if err != nil { + t.Fatal(err) + } + if err := db.AssignResourceToUser(user.ID, "conversation", conversation.ID); err != nil { + t.Fatal(err) + } + if _, err := db.CreateVulnerability(&database.Vulnerability{ConversationID: conversation.ID, Title: "finding", Severity: "high", Target: "192.0.2.42"}); err != nil { + t.Fatal(err) + } + scanCtx := mcp.WithMCPConversationID(ctx, conversation.ID) + result, _, err = server.CallTool(scanCtx, builtin.ToolCompleteAssetScan, map[string]interface{}{"id": id}) + if err != nil || result == nil || result.IsError { + t.Fatalf("complete scan result=%#v err=%v", result, err) + } + scanned, err := db.GetAsset(id, database.RBACListAccess{UserID: user.ID, Scope: database.RBACScopeAssigned}) + if err != nil || scanned.LastScanAt == nil || scanned.LastScanConversationID != conversation.ID || scanned.VulnerabilityCount != 1 { + t.Fatalf("scan fields not updated: %#v err=%v", scanned, err) + } + + result, _, err = server.CallTool(ctx, builtin.ToolDeleteAsset, map[string]interface{}{"id": id}) + if err != nil || result == nil || result.IsError { + t.Fatalf("delete asset result=%#v err=%v", result, err) + } + if _, err := db.GetAsset(id, database.RBACListAccess{Scope: database.RBACScopeAll}); err == nil { + t.Fatal("asset still exists after delete") + } +} + +func toolResultText(result *mcp.ToolResult) string { + var b strings.Builder + if result == nil { + return "" + } + for _, content := range result.Content { + b.WriteString(content.Text) + } + return b.String() +} + +func TestAssetReadToolsRespectConversationProjectScope(t *testing.T) { + db, err := database.NewDB(filepath.Join(t.TempDir(), "asset-project-scope.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + projectA, err := db.CreateProject(&database.Project{Name: "Project A"}) + if err != nil { + t.Fatal(err) + } + projectB, err := db.CreateProject(&database.Project{Name: "Project B"}) + if err != nil { + t.Fatal(err) + } + assets := []*database.Asset{ + {ProjectID: projectA.ID, IP: "192.0.2.10", Protocol: "https"}, + {ProjectID: projectB.ID, IP: "192.0.2.20", Protocol: "https"}, + {IP: "192.0.2.30", Protocol: "https"}, + } + if result, err := db.UpsertAssets(assets, "", true); err != nil || result.Created != len(assets) { + t.Fatalf("seed assets result=%#v err=%v", result, err) + } + + bound, err := db.CreateConversation("bound", database.ConversationCreateMeta{ProjectID: projectA.ID}) + if err != nil { + t.Fatal(err) + } + unbound, err := db.CreateConversation("unbound", database.ConversationCreateMeta{}) + if err != nil { + t.Fatal(err) + } + principal := authctx.NewPrincipal("admin", "admin", database.RBACScopeAll, map[string]bool{"asset:read": true}) + ctx := authctx.WithPrincipal(context.Background(), principal) + server := mcp.NewServer(zap.NewNop()) + server.SetToolAuthorizer(mcpToolAuthorizer(db)) + registerAssetTools(server, db, zap.NewNop()) + + boundCtx := mcp.WithMCPConversationID(ctx, bound.ID) + result, _, err := server.CallTool(boundCtx, builtin.ToolQueryAssets, map[string]interface{}{}) + text := toolResultText(result) + if err != nil || result == nil || result.IsError || !strings.Contains(text, assets[0].ID) || strings.Contains(text, assets[1].ID) || strings.Contains(text, assets[2].ID) { + t.Fatalf("bound query escaped project scope: result=%#v text=%q err=%v", result, text, err) + } + + // Even an explicit foreign project_id cannot override the conversation boundary. + result, _, err = server.CallTool(boundCtx, builtin.ToolQueryAssets, map[string]interface{}{"project_id": projectB.ID}) + text = toolResultText(result) + if err != nil || result == nil || result.IsError || !strings.Contains(text, assets[0].ID) || strings.Contains(text, assets[1].ID) { + t.Fatalf("project_id overrode conversation scope: result=%#v text=%q err=%v", result, text, err) + } + + result, _, err = server.CallTool(boundCtx, builtin.ToolGetAsset, map[string]interface{}{"id": assets[1].ID}) + if err != nil || result == nil || !result.IsError { + t.Fatalf("bound get read a foreign-project asset: result=%#v err=%v", result, err) + } + + unboundCtx := mcp.WithMCPConversationID(ctx, unbound.ID) + result, _, err = server.CallTool(unboundCtx, builtin.ToolQueryAssets, map[string]interface{}{"page_size": 10}) + text = toolResultText(result) + if err != nil || result == nil || result.IsError || !strings.Contains(text, assets[0].ID) || !strings.Contains(text, assets[1].ID) || !strings.Contains(text, assets[2].ID) { + t.Fatalf("unbound query did not retain all-assets behavior: result=%#v text=%q err=%v", result, text, err) + } +} diff --git a/internal/app/c2_hitl_bridge.go b/internal/app/c2_hitl_bridge.go new file mode 100644 index 00000000..7477d5a5 --- /dev/null +++ b/internal/app/c2_hitl_bridge.go @@ -0,0 +1,228 @@ +package app + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "strings" + "time" + + "cyberstrike-ai/internal/c2" + "cyberstrike-ai/internal/database" + + "github.com/google/uuid" + "go.uber.org/zap" +) + +// C2HITLBridge 实现 C2 Manager 的 HITLBridge 接口,将危险任务桥接到现有 HITL 审批流。 +// 审批记录写入 hitl_interrupts 表,与现有 HITL 系统共享前端审批 UI。 +type C2HITLBridge struct { + db *database.DB + logger *zap.Logger + timeout time.Duration + getConvID func() string +} + +// NewC2HITLBridge 创建 C2 HITL 桥 +func NewC2HITLBridge(db *database.DB, logger *zap.Logger) *C2HITLBridge { + return &C2HITLBridge{ + db: db, + logger: logger, + timeout: 5 * time.Minute, + getConvID: func() string { return "" }, + } +} + +// SetConversationIDGetter 设置获取当前对话 ID 的函数 +func (b *C2HITLBridge) SetConversationIDGetter(fn func() string) { + b.getConvID = fn +} + +// SetTimeout 设置审批超时(0 表示不超时) +func (b *C2HITLBridge) SetTimeout(d time.Duration) { + b.timeout = d +} + +// RequestApproval 实现 HITLBridge 接口:写入 hitl_interrupts 表并轮询等待审批结果 +func (b *C2HITLBridge) RequestApproval(ctx context.Context, req c2.HITLApprovalRequest) error { + interruptID := "hitl_c2_" + strings.ReplaceAll(uuid.New().String(), "-", "")[:14] + now := time.Now() + + convID := req.ConversationID + if convID == "" { + convID = b.getConvID() + } + if convID == "" { + convID = "c2_system" + } + + payload, _ := json.Marshal(map[string]interface{}{ + "task_id": req.TaskID, + "session_id": req.SessionID, + "task_type": req.TaskType, + "payload": req.PayloadJSON, + "source": req.Source, + "reason": req.Reason, + "c2_operation": true, + }) + + _, err := b.db.Exec(`INSERT INTO hitl_interrupts + (id, conversation_id, message_id, mode, tool_name, tool_call_id, payload, status, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, 'pending', ?)`, + interruptID, convID, "", "approval", + c2.MCPToolC2Task, req.TaskID, + string(payload), now, + ) + if err != nil { + b.logger.Error("C2 HITL: 创建审批记录失败,拒绝执行", zap.Error(err)) + return fmt.Errorf("C2 HITL 审批记录创建失败,安全起见拒绝执行: %w", err) + } + + b.logger.Info("C2 HITL: 等待人工审批", + zap.String("interrupt_id", interruptID), + zap.String("task_id", req.TaskID), + zap.String("task_type", req.TaskType), + ) + + // Poll DB waiting for decision + ticker := time.NewTicker(500 * time.Millisecond) + defer ticker.Stop() + + var deadline <-chan time.Time + if b.timeout > 0 { + timer := time.NewTimer(b.timeout) + defer timer.Stop() + deadline = timer.C + } + + for { + select { + case <-ctx.Done(): + _, _ = b.db.Exec(`UPDATE hitl_interrupts SET status='cancelled', decision='reject', + decision_comment='context cancelled', decided_at=? WHERE id=? AND status='pending'`, + time.Now(), interruptID) + return ctx.Err() + + case <-deadline: + _, _ = b.db.Exec(`UPDATE hitl_interrupts SET status='timeout', decision='reject', + decision_comment='C2 HITL timeout auto-reject for safety', decided_at=? WHERE id=? AND status='pending'`, + time.Now(), interruptID) + b.logger.Warn("C2 HITL: 审批超时,安全起见拒绝执行", zap.String("interrupt_id", interruptID)) + return fmt.Errorf("C2 HITL 审批超时,危险任务已被自动拒绝") + + case <-ticker.C: + var status, decision string + err := b.db.QueryRow(`SELECT status, COALESCE(decision, '') FROM hitl_interrupts WHERE id = ?`, + interruptID).Scan(&status, &decision) + if err != nil { + if err == sql.ErrNoRows { + return nil + } + continue + } + switch status { + case "decided", "timeout": + if decision == "reject" { + return fmt.Errorf("C2 危险任务被人工拒绝") + } + return nil + case "cancelled": + return fmt.Errorf("C2 审批已取消") + case "pending": + continue + default: + continue + } + } + } +} + +// C2HooksConfig 配置 C2 Manager 的 Hooks +type C2HooksConfig struct { + DB *database.DB + Logger *zap.Logger + AttackChainRecord func(session *database.C2Session, phase string, description string) + VulnRecord func(session *database.C2Session, title string, severity string) +} + +// SetupC2Hooks 设置 C2 Manager 的业务钩子 +func SetupC2Hooks(cfg *C2HooksConfig) c2.Hooks { + return c2.Hooks{ + OnSessionFirstSeen: func(session *database.C2Session) { + // 新会话上线 + cfg.Logger.Info("C2 Session first seen", + zap.String("session_id", session.ID), + zap.String("hostname", session.Hostname), + zap.String("os", session.OS), + zap.String("arch", session.Arch), + ) + + // 记录漏洞(初始访问点) + if cfg.VulnRecord != nil { + cfg.VulnRecord(session, fmt.Sprintf("C2 Session Established: %s@%s", session.Username, session.Hostname), "high") + } + + // 记录攻击链(Initial Access) + if cfg.AttackChainRecord != nil { + cfg.AttackChainRecord(session, "initial-access", fmt.Sprintf("Implant beacon from %s/%s", session.Hostname, session.InternalIP)) + } + }, + OnTaskCompleted: func(task *database.C2Task, sessionID string) { + // 任务完成 + cfg.Logger.Debug("C2 Task completed", + zap.String("task_id", task.ID), + zap.String("task_type", task.TaskType), + zap.String("status", task.Status), + ) + + // 根据任务类型记录攻击链 + if cfg.AttackChainRecord != nil { + session, _ := cfg.DB.GetC2Session(sessionID) + if session != nil { + phase := taskToAttackPhase(task.TaskType) + if phase != "" { + cfg.AttackChainRecord(session, phase, fmt.Sprintf("Task %s: %s", task.TaskType, task.Status)) + } + } + } + }, + } +} + +// taskToAttackPhase 将任务类型映射到 ATT&CK 阶段 +func taskToAttackPhase(taskType string) string { + switch taskType { + case "exec", "shell": + return "execution" + case "upload": + return "persistence" + case "download": + return "exfiltration" + case "screenshot": + return "collection" + case "kill_proc": + return "impact" + case "port_fwd", "socks_start": + return "lateral-movement" + case "load_assembly": + return "defense-evasion" + case "persist": + return "persistence" + case "self_delete": + return "defense-evasion" + default: + return "execution" + } +} + +// SetupC2HITLBridgeWithAgent 设置 HITL 桥接器 +// 这个函数将由 App 调用,注入必要的依赖 +func SetupC2HITLBridgeWithAgent(db *database.DB, logger *zap.Logger) c2.HITLBridge { + return &C2HITLBridge{ + db: db, + logger: logger, + timeout: 5 * time.Minute, + getConvID: func() string { return "" }, + } +} diff --git a/internal/app/c2_lifecycle.go b/internal/app/c2_lifecycle.go new file mode 100644 index 00000000..af651c39 --- /dev/null +++ b/internal/app/c2_lifecycle.go @@ -0,0 +1,104 @@ +package app + +import ( + "context" + + "cyberstrike-ai/internal/c2" + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/handler" + + "go.uber.org/zap" +) + +// setupC2Runtime 创建 C2 Manager、看门狗与取消函数;不注册 MCP 工具(由 Apply 统一 ClearTools 后注册)。 +func setupC2Runtime( + cfg *config.Config, + db *database.DB, + agentHandler *handler.AgentHandler, + logger *zap.Logger, +) (*c2.Manager, *c2.SessionWatchdog, context.CancelFunc) { + if !cfg.C2.EnabledEffective() { + return nil, nil, nil + } + c2Manager := c2.NewManager(db, logger, "tmp/c2") + c2Manager.Registry().Register(string(c2.ListenerTypeTCPReverse), c2.NewTCPReverseListener) + c2Manager.Registry().Register(string(c2.ListenerTypeHTTPBeacon), c2.NewHTTPBeaconListener) + c2Manager.Registry().Register(string(c2.ListenerTypeHTTPSBeacon), c2.NewHTTPSBeaconListener) + c2Manager.Registry().Register(string(c2.ListenerTypeWebSocket), c2.NewWebSocketListener) + c2HITLBridge := NewC2HITLBridge(db, logger) + c2Manager.SetHITLBridge(c2HITLBridge) + c2Manager.SetHITLDangerousGate(func(conversationID, toolName string) bool { + return agentHandler.HITLNeedsToolApproval(conversationID, toolName) + }) + c2Hooks := SetupC2Hooks(&C2HooksConfig{ + DB: db, + Logger: logger, + AttackChainRecord: func(session *database.C2Session, phase string, description string) { + logger.Info("C2 Attack Chain", + zap.String("session_id", session.ID), + zap.String("phase", phase), + zap.String("desc", description), + ) + }, + VulnRecord: func(session *database.C2Session, title string, severity string) { + logger.Info("C2 Vulnerability", + zap.String("session_id", session.ID), + zap.String("title", title), + zap.String("severity", severity), + ) + }, + }) + c2Manager.SetHooks(c2Hooks) + c2Manager.RestoreRunningListeners() + c2Watchdog := c2.NewSessionWatchdog(c2Manager) + watchdogCtx, watchdogCancel := context.WithCancel(context.Background()) + go c2Watchdog.Run(watchdogCtx) + return c2Manager, c2Watchdog, watchdogCancel +} + +// ReconcileC2AfterConfigApply 根据当前内存配置启停 C2(不写盘;在 Apply 中 ClearTools 之前调用)。 +func (a *App) ReconcileC2AfterConfigApply() error { + if !a.config.C2.EnabledEffective() { + a.shutdownC2() + return nil + } + if a.c2Manager != nil { + return nil + } + if a.db == nil || a.agentHandler == nil { + return nil + } + m, wd, cancel := setupC2Runtime(a.config, a.db, a.agentHandler, a.logger.Logger) + if m == nil { + return nil + } + a.c2Manager = m + a.c2Watchdog = wd + a.c2WatchdogCancel = cancel + if a.c2Handler != nil { + a.c2Handler.SetManager(m) + } + a.logger.Info("C2 子系统已按配置启动") + return nil +} + +// shutdownC2 停止看门狗与所有监听器,并断开 Handler 引用。 +func (a *App) shutdownC2() { + had := a.c2WatchdogCancel != nil || a.c2Manager != nil + if a.c2WatchdogCancel != nil { + a.c2WatchdogCancel() + a.c2WatchdogCancel = nil + } + a.c2Watchdog = nil + if a.c2Manager != nil { + a.c2Manager.Close() + a.c2Manager = nil + } + if a.c2Handler != nil { + a.c2Handler.SetManager(nil) + } + if had { + a.logger.Info("C2 子系统已关闭") + } +} diff --git a/internal/app/c2_tools.go b/internal/app/c2_tools.go new file mode 100644 index 00000000..6cca39a3 --- /dev/null +++ b/internal/app/c2_tools.go @@ -0,0 +1,919 @@ +package app + +import ( + "context" + "encoding/json" + "fmt" + "path/filepath" + "strconv" + "strings" + "time" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/authctx" + "cyberstrike-ai/internal/c2" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" + + "github.com/google/uuid" + "go.uber.org/zap" +) + +// registerC2Tools 注册所有 C2 MCP 工具(合并同类项,减少工具数量以节省上下文 token)。 +// webListenPort 为本进程 Web/API 监听端口(配置 server.port,启动时已加载),用于 MCP 描述中提示勿与 C2 bind_port 冲突。 +func registerC2Tools(mcpServer *mcp.Server, c2Manager *c2.Manager, logger *zap.Logger, webListenPort int) { + registerC2ListenerTool(mcpServer, c2Manager, logger, webListenPort) + registerC2SessionTool(mcpServer, c2Manager, logger) + registerC2TaskTool(mcpServer, c2Manager, logger) + registerC2TaskManageTool(mcpServer, c2Manager, logger) + registerC2PayloadTool(mcpServer, c2Manager, logger, webListenPort) + registerC2EventTool(mcpServer, c2Manager, logger) + registerC2ProfileTool(mcpServer, c2Manager, logger) + registerC2FileTool(mcpServer, c2Manager, logger) + logger.Debug("C2 MCP tools registered (8 unified tools)") +} + +func makeC2Result(data interface{}, err error) (*mcp.ToolResult, error) { + if err != nil { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: err.Error()}}, + IsError: true, + }, nil + } + text, _ := json.Marshal(data) + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: string(text)}}, + }, nil +} + +// ============================================================================ +// c2_listener — 监听器统一工具 +// ============================================================================ + +func registerC2ListenerTool(s *mcp.Server, m *c2.Manager, l *zap.Logger, webListenPort int) { + s.RegisterTool(mcp.Tool{ + Name: builtin.ToolC2Listener, + Description: fmt.Sprintf(`C2 监听器管理。通过 action 参数选择操作: +- list: 列出所有监听器 +- get: 获取监听器详情(需 listener_id) +- create: 创建监听器(需 name, type, bind_port)。成功时除 listener 外会返回 implant_token(仅此一次,用于 X-Implant-Token / oneliner;list/get/start 不再返回) +- update: 更新监听器配置(需 listener_id,可改 name/bind_host/bind_port/remark/config/callback_host) +- start: 启动监听器(需 listener_id) +- stop: 停止监听器(需 listener_id) +- delete: 删除监听器(需 listener_id) +监听器类型: tcp_reverse, http_beacon, https_beacon, websocket +tcp_reverse 默认仅接受 CSB1 加密 Beacon(AES-GCM + ImplantToken)才登记会话;经典 bash/nc 反弹需在 config.allow_legacy_shell=true(公网不推荐)。 +端口约束:create/update 的 bind_port 禁止与本平台 Web/API 所用端口相同。当前本服务该端口为 %d(配置项 server.port,随进程启动从配置文件加载)。若 bind_port 与此相同会导致本服务或监听器 bind 失败、Beacon/oneliner 误连到 Web 而非 C2。请为监听器另选空闲端口。`, webListenPort), + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "action": map[string]interface{}{"type": "string", "description": "操作: list/get/create/update/start/stop/delete", "enum": []string{"list", "get", "create", "update", "start", "stop", "delete"}}, + "listener_id": map[string]interface{}{"type": "string", "description": "监听器 ID(get/update/start/stop/delete 需要)"}, + "name": map[string]interface{}{"type": "string", "description": "监听器名称(create/update)"}, + "type": map[string]interface{}{"type": "string", "description": "监听器类型(create)", "enum": []string{"tcp_reverse", "http_beacon", "https_beacon", "websocket"}}, + "bind_host": map[string]interface{}{"type": "string", "description": "绑定地址,默认 127.0.0.1;外网监听常用 0.0.0.0"}, + "callback_host": map[string]interface{}{"type": "string", "description": "可选:植入端/Payload 回连主机名(公网 IP 或域名)。写入 config_json;生成 oneliner/beacon 时优先于 bind_host。update 时传入空字符串可清除"}, + "bind_port": map[string]interface{}{"type": "integer", "description": fmt.Sprintf("绑定端口(create 必填)。须 ≠ %d(当前本服务 Web/API 端口,配置 server.port)", webListenPort), "minimum": 1, "maximum": 65535}, + "project_id": map[string]interface{}{"type": "string", "description": "所属项目 ID。create 省略时默认使用当前对话绑定项目;未绑定项目的对话则创建未绑定监听器"}, + "profile_id": map[string]interface{}{"type": "string", "description": "Malleable Profile ID"}, + "remark": map[string]interface{}{"type": "string", "description": "备注"}, + "config": map[string]interface{}{"type": "object", "description": "高级配置(beacon 路径/TLS/OPSEC 等),create/update 可用。tcp_reverse 可选 allow_legacy_shell:true 允许未加密经典 shell(默认 false)"}, + }, + "required": []string{"action"}, + }, + }, func(ctx context.Context, params map[string]interface{}) (*mcp.ToolResult, error) { + action := getString(params, "action") + id := getString(params, "listener_id") + + switch action { + case "list": + listeners, err := m.DB().ListC2ListenersForAccess(c2ToolAccess(ctx), mcpEffectiveProjectFilter(ctx, m.DB())) + if err != nil { + return makeC2Result(nil, err) + } + for _, li := range listeners { + li.EncryptionKey = "" + li.ImplantToken = "" + } + return makeC2Result(map[string]interface{}{"listeners": listeners, "count": len(listeners)}, nil) + + case "get": + listener, err := m.DB().GetC2Listener(id) + if err != nil { + return makeC2Result(nil, err) + } + if listener == nil { + return makeC2Result(nil, fmt.Errorf("listener not found")) + } + listener.EncryptionKey = "" + listener.ImplantToken = "" + return makeC2Result(map[string]interface{}{"listener": listener}, nil) + + case "create": + var cfg *c2.ListenerConfig + if cfgRaw, ok := params["config"]; ok && cfgRaw != nil { + cfgBytes, _ := json.Marshal(cfgRaw) + cfg = &c2.ListenerConfig{} + _ = json.Unmarshal(cfgBytes, cfg) + } + projectID := strings.TrimSpace(getString(params, "project_id")) + if projectID == "" { + projectID = mcpEffectiveProjectFilter(ctx, m.DB()) + if projectID == database.ProjectFilterUnbound { + projectID = "" + } + } + input := c2.CreateListenerInput{ + Name: getString(params, "name"), + Type: getString(params, "type"), + BindHost: getString(params, "bind_host"), + BindPort: int(getFloat64(params, "bind_port")), + ProfileID: getString(params, "profile_id"), + Remark: getString(params, "remark"), + ProjectID: projectID, + Config: cfg, + CallbackHost: getString(params, "callback_host"), + } + listener, err := m.CreateListener(input) + if err != nil { + return makeC2Result(nil, err) + } + if principal, ok := authctx.PrincipalFromContext(ctx); ok { + _ = m.DB().SetResourceOwner("c2_listener", listener.ID, principal.UserID) + _ = m.DB().AssignResourceToUser(principal.UserID, "c2_listener", listener.ID) + } + implantToken := listener.ImplantToken + listener.EncryptionKey = "" + listener.ImplantToken = "" + return makeC2Result(map[string]interface{}{ + "listener": listener, + "implant_token": implantToken, + }, nil) + + case "update": + listener, err := m.DB().GetC2Listener(id) + if err != nil { + return makeC2Result(nil, err) + } + if listener == nil { + return makeC2Result(nil, fmt.Errorf("listener not found")) + } + if m.IsListenerRunning(id) { + newHost := getString(params, "bind_host") + newPort := int(getFloat64(params, "bind_port")) + if (newHost != "" && newHost != listener.BindHost) || (newPort > 0 && newPort != listener.BindPort) { + return makeC2Result(nil, fmt.Errorf("cannot modify bind address while listener is running")) + } + } + if v := getString(params, "name"); v != "" { + listener.Name = v + } + if v := getString(params, "bind_host"); v != "" { + listener.BindHost = v + } + if v := int(getFloat64(params, "bind_port")); v > 0 { + listener.BindPort = v + } + if v := getString(params, "profile_id"); v != "" { + listener.ProfileID = v + } + if v, ok := params["remark"]; ok { + listener.Remark, _ = v.(string) + } + if cfgRaw, ok := params["config"]; ok && cfgRaw != nil { + cfgBytes, _ := json.Marshal(cfgRaw) + listener.ConfigJSON = string(cfgBytes) + } + if _, ok := params["callback_host"]; ok { + pcfg := &c2.ListenerConfig{} + raw := strings.TrimSpace(listener.ConfigJSON) + if raw == "" { + raw = "{}" + } + _ = json.Unmarshal([]byte(raw), pcfg) + pcfg.CallbackHost = strings.TrimSpace(getString(params, "callback_host")) + pcfg.ApplyDefaults() + cfgBytes, err := json.Marshal(pcfg) + if err != nil { + return makeC2Result(nil, err) + } + listener.ConfigJSON = string(cfgBytes) + } + if err := m.DB().UpdateC2Listener(listener); err != nil { + return makeC2Result(nil, err) + } + listener.EncryptionKey = "" + listener.ImplantToken = "" + return makeC2Result(map[string]interface{}{"listener": listener}, nil) + + case "start": + listener, err := m.StartListener(id) + if err != nil { + return makeC2Result(nil, err) + } + listener.EncryptionKey = "" + listener.ImplantToken = "" + return makeC2Result(map[string]interface{}{"listener": listener}, nil) + + case "stop": + err := m.StopListener(id) + return makeC2Result(map[string]interface{}{"stopped": err == nil}, err) + + case "delete": + err := m.DeleteListener(id) + return makeC2Result(map[string]interface{}{"deleted": err == nil}, err) + + default: + return makeC2Result(nil, fmt.Errorf("unknown action: %s", action)) + } + }) +} + +// ============================================================================ +// c2_session — 会话统一工具 +// ============================================================================ + +func registerC2SessionTool(s *mcp.Server, m *c2.Manager, l *zap.Logger) { + s.RegisterTool(mcp.Tool{ + Name: builtin.ToolC2Session, + Description: `C2 会话管理。通过 action 参数选择操作: +- list: 列出会话(可按 listener_id/status/os/search/suspicious 过滤) +- get: 获取会话详情及最近任务历史(需 session_id) +- set_sleep: 设置心跳间隔(需 session_id) +- kill: 下发 exit 任务让 implant 退出(需 session_id) +- delete: 删除单个会话记录(需 session_id) +- delete_batch: 批量删除会话(需 session_ids 数组)`, + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "action": map[string]interface{}{"type": "string", "description": "操作: list/get/set_sleep/kill/delete/delete_batch", "enum": []string{"list", "get", "set_sleep", "kill", "delete", "delete_batch"}}, + "session_id": map[string]interface{}{"type": "string", "description": "会话 ID(get/set_sleep/kill/delete 需要)"}, + "session_ids": map[string]interface{}{"type": "array", "items": map[string]interface{}{"type": "string"}, "description": "会话 ID 列表(delete_batch)"}, + "listener_id": map[string]interface{}{"type": "string", "description": "按监听器过滤(list)"}, + "status": map[string]interface{}{"type": "string", "description": "按状态过滤: active/sleeping/dead/killed(list)"}, + "os": map[string]interface{}{"type": "string", "description": "按 OS 过滤: linux/windows/darwin(list)"}, + "search": map[string]interface{}{"type": "string", "description": "模糊搜索 hostname/username/IP(list)"}, + "suspicious": map[string]interface{}{"type": "boolean", "description": "仅疑似误报:离线且 tcp_* / unknown / PID 0(list)"}, + "limit": map[string]interface{}{"type": "integer", "description": "返回数量上限(list)"}, + "sleep_seconds": map[string]interface{}{"type": "integer", "description": "心跳间隔秒数(set_sleep)"}, + "jitter_percent": map[string]interface{}{"type": "integer", "description": "抖动百分比 0-100(set_sleep)"}, + }, + "required": []string{"action"}, + }, + }, func(ctx context.Context, params map[string]interface{}) (*mcp.ToolResult, error) { + action := getString(params, "action") + id := getString(params, "session_id") + + switch action { + case "list": + filter := database.ListC2SessionsFilter{ + ListenerID: getString(params, "listener_id"), + ProjectID: mcpEffectiveProjectFilter(ctx, m.DB()), + Status: getString(params, "status"), + OS: getString(params, "os"), + Search: getString(params, "search"), + } + if limit := int(getFloat64(params, "limit")); limit > 0 { + filter.Limit = limit + } + if v, ok := params["suspicious"].(bool); ok && v { + filter.Suspicious = true + } + sessions, err := m.DB().ListC2SessionsForAccess(filter, c2ToolAccess(ctx)) + return makeC2Result(map[string]interface{}{"sessions": sessions, "count": len(sessions)}, err) + + case "get": + session, err := m.DB().GetC2Session(id) + if err != nil { + return makeC2Result(nil, err) + } + if session == nil { + return makeC2Result(nil, fmt.Errorf("session not found")) + } + tasks, _ := m.DB().ListC2Tasks(database.ListC2TasksFilter{SessionID: id, Limit: 10}) + return makeC2Result(map[string]interface{}{"session": session, "tasks": tasks}, nil) + + case "set_sleep": + sleep := int(getFloat64(params, "sleep_seconds")) + jitter := int(getFloat64(params, "jitter_percent")) + task, err := m.SetSessionSleep(id, sleep, jitter) + out := map[string]interface{}{ + "updated": err == nil, + "sleep_seconds": sleep, + "jitter_percent": jitter, + } + if task != nil { + out["task_id"] = task.ID + } + return makeC2Result(out, err) + + case "kill": + task, err := m.EnqueueTask(c2.EnqueueTaskInput{ + SessionID: id, + TaskType: c2.TaskTypeExit, + Payload: map[string]interface{}{}, + Source: "ai", + ConversationID: agent.ConversationIDFromContext(ctx), + UserCtx: ctx, + }) + return makeC2Result(map[string]interface{}{"task": task}, err) + + case "delete": + err := m.DB().DeleteC2Session(id) + return makeC2Result(map[string]interface{}{"deleted": err == nil}, err) + + case "delete_batch": + rawIDs, _ := params["session_ids"].([]interface{}) + ids := make([]string, 0, len(rawIDs)) + for _, v := range rawIDs { + if s, ok := v.(string); ok && strings.TrimSpace(s) != "" { + ids = append(ids, strings.TrimSpace(s)) + } + } + n, err := m.DB().DeleteC2SessionsByIDs(ids) + return makeC2Result(map[string]interface{}{"deleted": n}, err) + + default: + return makeC2Result(nil, fmt.Errorf("unknown action: %s", action)) + } + }) +} + +// ============================================================================ +// c2_task — 任务下发统一工具(合并所有 task 类型) +// ============================================================================ + +func registerC2TaskTool(s *mcp.Server, m *c2.Manager, l *zap.Logger) { + s.RegisterTool(mcp.Tool{ + Name: builtin.ToolC2Task, + Description: `在 C2 会话上下发任务。所有任务类型通过 task_type 参数指定: +- exec: 执行命令(需 command) +- shell: 交互式命令,保持 cwd(需 command) +- pwd/ps/screenshot/socks_stop: 无额外参数 +- cd/ls: 需 path +- kill_proc: 需 pid +- upload: 需 remote_path + file_id +- download: 需 remote_path +- port_fwd: 需 action(start/stop) + local_port + remote_host + remote_port +- socks_start: 需 port(默认 1080) +- load_assembly: 需 data(base64) 或 file_id,可选 args +- persist: 可选 method(auto/cron/bashrc/launchagent/registry/schtasks) +返回 task_id,用 c2_task_manage 的 wait/get_result 获取结果。`, + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "session_id": map[string]interface{}{"type": "string", "description": "C2 会话 ID(s_xxx)"}, + "task_type": map[string]interface{}{"type": "string", "description": "任务类型", "enum": []string{"exec", "shell", "pwd", "cd", "ls", "ps", "kill_proc", "upload", "download", "screenshot", "port_fwd", "socks_start", "socks_stop", "load_assembly", "persist"}}, + "command": map[string]interface{}{"type": "string", "description": "命令(exec/shell)"}, + "path": map[string]interface{}{"type": "string", "description": "路径(cd/ls)"}, + "pid": map[string]interface{}{"type": "integer", "description": "进程 ID(kill_proc)"}, + "remote_path": map[string]interface{}{"type": "string", "description": "远程路径(upload/download)"}, + "file_id": map[string]interface{}{"type": "string", "description": "服务端文件 ID(upload/load_assembly)"}, + "data": map[string]interface{}{"type": "string", "description": "base64 数据(load_assembly)"}, + "args": map[string]interface{}{"type": "string", "description": "命令行参数(load_assembly)"}, + "action": map[string]interface{}{"type": "string", "description": "start/stop(port_fwd)"}, + "local_port": map[string]interface{}{"type": "integer", "description": "本地端口(port_fwd)"}, + "remote_host": map[string]interface{}{"type": "string", "description": "远程主机(port_fwd)"}, + "remote_port": map[string]interface{}{"type": "integer", "description": "远程端口(port_fwd)"}, + "port": map[string]interface{}{"type": "integer", "description": "SOCKS5 端口(socks_start),默认 1080"}, + "method": map[string]interface{}{"type": "string", "description": "持久化方法(persist): auto/cron/bashrc/launchagent/registry/schtasks"}, + "timeout_seconds": map[string]interface{}{"type": "integer", "description": "超时秒数,默认 60"}, + }, + "required": []string{"session_id", "task_type"}, + }, + }, func(ctx context.Context, params map[string]interface{}) (*mcp.ToolResult, error) { + sessionID := getString(params, "session_id") + taskTypeStr := getString(params, "task_type") + taskType := c2.TaskType(taskTypeStr) + timeout := getFloat64(params, "timeout_seconds") + + payload := map[string]interface{}{"timeout_seconds": timeout} + + switch taskType { + case c2.TaskTypeExec, c2.TaskTypeShell: + payload["command"] = getString(params, "command") + case c2.TaskTypeCd, c2.TaskTypeLs: + payload["path"] = getString(params, "path") + case c2.TaskTypeKillProc: + payload["pid"] = params["pid"] + case c2.TaskTypeUpload: + payload["remote_path"] = getString(params, "remote_path") + payload["file_id"] = getString(params, "file_id") + case c2.TaskTypeDownload: + payload["remote_path"] = getString(params, "remote_path") + case c2.TaskTypePortFwd: + payload["action"] = getString(params, "action") + payload["local_port"] = params["local_port"] + payload["remote_host"] = getString(params, "remote_host") + payload["remote_port"] = params["remote_port"] + case c2.TaskTypeSocksStart: + payload["port"] = params["port"] + case c2.TaskTypeLoadAssembly: + payload["data"] = getString(params, "data") + payload["file_id"] = getString(params, "file_id") + payload["args"] = getString(params, "args") + case c2.TaskTypePersist: + payload["method"] = getString(params, "method") + case c2.TaskTypePwd, c2.TaskTypePs, c2.TaskTypeScreenshot, c2.TaskTypeSocksStop: + // no extra params + default: + return makeC2Result(nil, fmt.Errorf("unsupported task_type: %s", taskTypeStr)) + } + + input := c2.EnqueueTaskInput{ + SessionID: sessionID, + TaskType: taskType, + Payload: payload, + Source: "ai", + ConversationID: agent.ConversationIDFromContext(ctx), + UserCtx: ctx, + } + task, err := m.EnqueueTask(input) + if err != nil { + return makeC2Result(nil, err) + } + return makeC2Result(map[string]interface{}{"task_id": task.ID, "status": task.Status}, nil) + }) +} + +// ============================================================================ +// c2_task_manage — 任务管理工具(查询/等待/取消) +// ============================================================================ + +func registerC2TaskManageTool(s *mcp.Server, m *c2.Manager, l *zap.Logger) { + s.RegisterTool(mcp.Tool{ + Name: builtin.ToolC2TaskManage, + Description: `C2 任务管理。通过 action 参数选择操作: +- get_result: 获取任务详情和结果(需 task_id) +- wait: 阻塞等待任务完成并返回结果(需 task_id) +- list: 列出任务(可按 session_id/status 过滤) +- cancel: 取消排队中的任务(需 task_id)`, + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "action": map[string]interface{}{"type": "string", "description": "操作: get_result/wait/list/cancel", "enum": []string{"get_result", "wait", "list", "cancel"}}, + "task_id": map[string]interface{}{"type": "string", "description": "任务 ID(get_result/wait/cancel 需要)"}, + "session_id": map[string]interface{}{"type": "string", "description": "按会话过滤(list)"}, + "status": map[string]interface{}{"type": "string", "description": "按状态过滤: queued/sent/running/success/failed/cancelled(list)"}, + "limit": map[string]interface{}{"type": "integer", "description": "返回数量上限(list)"}, + "timeout_seconds": map[string]interface{}{"type": "integer", "description": "等待超时秒数(wait),默认 60"}, + }, + "required": []string{"action"}, + }, + }, func(ctx context.Context, params map[string]interface{}) (*mcp.ToolResult, error) { + action := getString(params, "action") + + switch action { + case "get_result": + id := getString(params, "task_id") + task, err := m.DB().GetC2Task(id) + if err != nil { + return makeC2Result(nil, err) + } + if task == nil { + return makeC2Result(nil, fmt.Errorf("task not found")) + } + return makeC2Result(map[string]interface{}{"task": task}, nil) + + case "wait": + id := getString(params, "task_id") + timeout := int(getFloat64(params, "timeout_seconds")) + if timeout <= 0 { + timeout = 60 + } + deadline := time.Now().Add(time.Duration(timeout) * time.Second) + for time.Now().Before(deadline) { + task, err := m.DB().GetC2Task(id) + if err != nil { + return makeC2Result(nil, err) + } + if task == nil { + return makeC2Result(nil, fmt.Errorf("task not found")) + } + if task.Status == "success" || task.Status == "failed" || task.Status == "cancelled" { + return makeC2Result(map[string]interface{}{"task": task}, nil) + } + select { + case <-time.After(500 * time.Millisecond): + case <-ctx.Done(): + return makeC2Result(nil, ctx.Err()) + } + } + return makeC2Result(nil, fmt.Errorf("timeout waiting for task completion")) + + case "list": + filter := database.ListC2TasksFilter{ + SessionID: getString(params, "session_id"), + ProjectID: mcpEffectiveProjectFilter(ctx, m.DB()), + Status: getString(params, "status"), + } + if limit := int(getFloat64(params, "limit")); limit > 0 { + filter.Limit = limit + } + tasks, err := m.DB().ListC2TasksForAccess(filter, c2ToolAccess(ctx)) + return makeC2Result(map[string]interface{}{"tasks": tasks, "count": len(tasks)}, err) + + case "cancel": + id := getString(params, "task_id") + err := m.CancelTask(id) + return makeC2Result(map[string]interface{}{"cancelled": err == nil}, err) + + default: + return makeC2Result(nil, fmt.Errorf("unknown action: %s", action)) + } + }) +} + +// ============================================================================ +// c2_payload — Payload 统一工具 +// ============================================================================ + +func registerC2PayloadTool(s *mcp.Server, m *c2.Manager, l *zap.Logger, webListenPort int) { + s.RegisterTool(mcp.Tool{ + Name: builtin.ToolC2Payload, + Description: fmt.Sprintf(`C2 Payload 生成。通过 action 参数选择操作: +- oneliner: 生成单行 payload。kind 必须与监听器协议一致,否则会失败: + • tcp_reverse:默认仅支持 build 加密 Beacon;若监听器 config.allow_legacy_shell=true,才可用 kind: bash, nc, nc_mkfifo, python, perl, powershell。 + • http_beacon / https_beacon / websocket:仅 HTTP(S) Beacon 轮询,oneliner 只能用 kind: curl_beacon(脚本内用 bash+curl,与「tcp 的 bash」不同)。curl_beacon 返回串末尾含「 &」用于把整个 bash -c 放后台;若用 exec/execute 同步执行,必须整段原样复制(含末尾 &)。若删掉 &,内部 while 死循环占满前台,调用会一直阻塞到超时/杀进程。 + • 公网部署 tcp_reverse 请用 build 生成加密 Beacon,勿开启 allow_legacy_shell。 + • 省略 kind 时,会按监听器类型自动选第一个兼容类型(HTTP 系默认为 curl_beacon)。 +- build: 交叉编译 beacon 二进制。支持 http_beacon / https_beacon / websocket / tcp_reverse(tcp_reverse 植入端回连后先发魔数 CSB1,再经 AES-GCM 解密且校验 ImplantToken 后才登记会话)。 +依赖的监听器 bind_port 须避开本服务 Web 端口 %d(配置 server.port,与 c2_listener 描述一致),否则 Beacon 无法正确回连。`, webListenPort), + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "action": map[string]interface{}{"type": "string", "description": "操作: oneliner/build", "enum": []string{"oneliner", "build"}}, + "listener_id": map[string]interface{}{"type": "string", "description": "监听器 ID(必填)。oneliner 前请确认该监听器的 type,再选兼容的 kind"}, + "kind": map[string]interface{}{"type": "string", "description": "仅 action=oneliner 需要。tcp_reverse: bash|nc|nc_mkfifo|python|perl|powershell;http_beacon|https_beacon|websocket: 仅 curl_beacon"}, + "host": map[string]interface{}{"type": "string", "description": "oneliner/build 可选覆盖:非空则强制用作植入回连主机。留空时顺序为:监听器 callback_host(create/update 的 callback_host 参数写入)→ bind_host(0.0.0.0 时尝试本机对外 IP 探测)"}, + "os": map[string]interface{}{"type": "string", "description": "目标 OS(build): linux/windows/darwin", "default": "linux"}, + "arch": map[string]interface{}{"type": "string", "description": "目标架构(build): amd64/arm64/386/arm", "default": "amd64"}, + "sleep_seconds": map[string]interface{}{"type": "integer", "description": "默认心跳间隔(build)"}, + "jitter_percent": map[string]interface{}{"type": "integer", "description": "默认抖动百分比(build)"}, + }, + "required": []string{"action", "listener_id"}, + }, + }, func(ctx context.Context, params map[string]interface{}) (*mcp.ToolResult, error) { + action := getString(params, "action") + listenerID := getString(params, "listener_id") + + switch action { + case "oneliner": + listener, err := m.DB().GetC2Listener(listenerID) + if err != nil { + return makeC2Result(nil, err) + } + if listener == nil { + return makeC2Result(nil, fmt.Errorf("listener not found")) + } + host := c2.ResolveBeaconDialHost(listener, getString(params, "host"), l, listenerID) + kind := c2.OnelinerKind(getString(params, "kind")) + if kind == "" { + compatible := c2.OnelinerKindsForListener(listener.Type) + if len(compatible) > 0 { + kind = compatible[0] + } + } + if !c2.IsOnelinerCompatible(listener.Type, kind) { + compatible := c2.OnelinerKindsForListener(listener.Type) + names := make([]string, len(compatible)) + for i, k := range compatible { + names[i] = string(k) + } + return makeC2Result(nil, fmt.Errorf("监听器类型 %s 不支持 %s,兼容类型: %v", listener.Type, kind, names)) + } + if err := c2.ValidateOnelinerForListener(listener, kind); err != nil { + return makeC2Result(nil, err) + } + input := c2.OnelinerInput{ + Kind: kind, + Host: host, + Port: listener.BindPort, + HTTPBaseURL: fmt.Sprintf("http://%s:%d", host, listener.BindPort), + ImplantToken: listener.ImplantToken, + } + oneliner, err := c2.GenerateOneliner(input) + if err != nil { + return makeC2Result(nil, err) + } + out := map[string]interface{}{ + "oneliner": oneliner, "kind": input.Kind, "host": host, "port": listener.BindPort, + } + if kind == c2.OnelinerCurl { + out["usage_note"] = "同步 exec/execute:整段原样执行(末尾须有「 &」)。去掉则 while 永不结束,工具会一直卡住。" + } + return makeC2Result(out, nil) + + case "build": + builder := c2.NewPayloadBuilder(m, l, "", "") + input := c2.PayloadBuilderInput{ + ListenerID: listenerID, + OS: getString(params, "os"), + Arch: getString(params, "arch"), + SleepSeconds: int(getFloat64(params, "sleep_seconds")), + JitterPercent: int(getFloat64(params, "jitter_percent")), + Host: strings.TrimSpace(getString(params, "host")), + } + result, err := builder.BuildBeacon(input) + if err != nil { + return makeC2Result(nil, err) + } + if principal, ok := authctx.PrincipalFromContext(ctx); ok { + _ = m.DB().RecordC2PayloadArtifact(filepath.Base(result.OutputPath), result.PayloadID, result.ListenerID, principal.UserID) + } + return makeC2Result(map[string]interface{}{ + "payload_id": result.PayloadID, "download_path": result.DownloadPath, + "os": result.OS, "arch": result.Arch, "size_bytes": result.SizeBytes, + }, nil) + + default: + return makeC2Result(nil, fmt.Errorf("unknown action: %s", action)) + } + }) +} + +// ============================================================================ +// c2_event — 事件查询工具 +// ============================================================================ + +func registerC2EventTool(s *mcp.Server, m *c2.Manager, l *zap.Logger) { + s.RegisterTool(mcp.Tool{ + Name: builtin.ToolC2Event, + Description: "获取 C2 事件(上线/掉线/任务/错误),支持按级别/类别/会话/任务/时间过滤", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "level": map[string]interface{}{"type": "string", "description": "级别过滤: info/warn/critical"}, + "category": map[string]interface{}{"type": "string", "description": "类别过滤: listener/session/task/payload/opsec"}, + "session_id": map[string]interface{}{"type": "string", "description": "按会话过滤"}, + "task_id": map[string]interface{}{"type": "string", "description": "按任务过滤"}, + "since": map[string]interface{}{"type": "string", "description": "起始时间(RFC3339 格式,如 2025-01-01T00:00:00Z)"}, + "limit": map[string]interface{}{"type": "integer", "default": 50, "description": "返回数量"}, + }, + }, + }, func(ctx context.Context, params map[string]interface{}) (*mcp.ToolResult, error) { + filter := database.ListC2EventsFilter{ + Level: getString(params, "level"), + Category: getString(params, "category"), + ProjectID: mcpEffectiveProjectFilter(ctx, m.DB()), + SessionID: getString(params, "session_id"), + TaskID: getString(params, "task_id"), + Limit: int(getFloat64(params, "limit")), + } + if filter.Limit <= 0 { + filter.Limit = 50 + } + if since := getString(params, "since"); since != "" { + if t, err := time.Parse(time.RFC3339, since); err == nil { + filter.Since = &t + } + } + events, err := m.DB().ListC2EventsForAccess(filter, c2ToolAccess(ctx)) + return makeC2Result(map[string]interface{}{"events": events, "count": len(events)}, err) + }) +} + +func c2ToolAccess(ctx context.Context) database.RBACListAccess { + principal, ok := authctx.PrincipalFromContext(ctx) + if !ok { + return database.RBACListAccess{Scope: database.RBACScopeAssigned} + } + return database.RBACListAccess{UserID: principal.UserID, Scope: principal.ScopeFor("c2:read")} +} + +// ============================================================================ +// c2_profile — Malleable Profile 管理工具(新增) +// ============================================================================ + +func registerC2ProfileTool(s *mcp.Server, m *c2.Manager, l *zap.Logger) { + s.RegisterTool(mcp.Tool{ + Name: builtin.ToolC2Profile, + Description: `C2 Malleable Profile 管理(控制 beacon 通信伪装)。通过 action 参数选择操作: +- list: 列出所有 Profile +- get: 获取 Profile 详情(需 profile_id) +- create: 创建 Profile(需 name,可选 user_agent/uris/request_headers/response_headers/body_template/jitter_min_ms/jitter_max_ms) +- update: 更新 Profile(需 profile_id) +- delete: 删除 Profile(需 profile_id)`, + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "action": map[string]interface{}{"type": "string", "description": "操作: list/get/create/update/delete", "enum": []string{"list", "get", "create", "update", "delete"}}, + "profile_id": map[string]interface{}{"type": "string", "description": "Profile ID(get/update/delete 需要)"}, + "name": map[string]interface{}{"type": "string", "description": "Profile 名称"}, + "user_agent": map[string]interface{}{"type": "string", "description": "User-Agent 字符串"}, + "uris": map[string]interface{}{"type": "array", "items": map[string]interface{}{"type": "string"}, "description": "beacon 请求的 URI 列表"}, + "request_headers": map[string]interface{}{"type": "object", "description": "自定义请求头"}, + "response_headers": map[string]interface{}{"type": "object", "description": "自定义响应头"}, + "body_template": map[string]interface{}{"type": "string", "description": "响应体模板"}, + "jitter_min_ms": map[string]interface{}{"type": "integer", "description": "最小抖动(毫秒)"}, + "jitter_max_ms": map[string]interface{}{"type": "integer", "description": "最大抖动(毫秒)"}, + }, + "required": []string{"action"}, + }, + }, func(ctx context.Context, params map[string]interface{}) (*mcp.ToolResult, error) { + action := getString(params, "action") + id := getString(params, "profile_id") + + switch action { + case "list": + profiles, err := m.DB().ListC2Profiles() + return makeC2Result(map[string]interface{}{"profiles": profiles, "count": len(profiles)}, err) + + case "get": + profile, err := m.DB().GetC2Profile(id) + if err != nil { + return makeC2Result(nil, err) + } + if profile == nil { + return makeC2Result(nil, fmt.Errorf("profile not found")) + } + return makeC2Result(map[string]interface{}{"profile": profile}, nil) + + case "create": + profile := &database.C2Profile{ + ID: "p_" + strings.ReplaceAll(uuid.New().String(), "-", "")[:14], + Name: getString(params, "name"), + UserAgent: getString(params, "user_agent"), + BodyTemplate: getString(params, "body_template"), + JitterMinMS: int(getFloat64(params, "jitter_min_ms")), + JitterMaxMS: int(getFloat64(params, "jitter_max_ms")), + CreatedAt: time.Now(), + } + if uris, ok := params["uris"]; ok { + if arr, ok := uris.([]interface{}); ok { + for _, u := range arr { + if s, ok := u.(string); ok { + profile.URIs = append(profile.URIs, s) + } + } + } + } + if rh, ok := params["request_headers"]; ok { + if m, ok := rh.(map[string]interface{}); ok { + profile.RequestHeaders = make(map[string]string) + for k, v := range m { + profile.RequestHeaders[k], _ = v.(string) + } + } + } + if rh, ok := params["response_headers"]; ok { + if m, ok := rh.(map[string]interface{}); ok { + profile.ResponseHeaders = make(map[string]string) + for k, v := range m { + profile.ResponseHeaders[k], _ = v.(string) + } + } + } + if err := m.DB().CreateC2Profile(profile); err != nil { + return makeC2Result(nil, err) + } + return makeC2Result(map[string]interface{}{"profile": profile}, nil) + + case "update": + profile, err := m.DB().GetC2Profile(id) + if err != nil { + return makeC2Result(nil, err) + } + if profile == nil { + return makeC2Result(nil, fmt.Errorf("profile not found")) + } + if v := getString(params, "name"); v != "" { + profile.Name = v + } + if v := getString(params, "user_agent"); v != "" { + profile.UserAgent = v + } + if v := getString(params, "body_template"); v != "" { + profile.BodyTemplate = v + } + if v := int(getFloat64(params, "jitter_min_ms")); v > 0 { + profile.JitterMinMS = v + } + if v := int(getFloat64(params, "jitter_max_ms")); v > 0 { + profile.JitterMaxMS = v + } + if uris, ok := params["uris"]; ok { + if arr, ok := uris.([]interface{}); ok { + profile.URIs = nil + for _, u := range arr { + if s, ok := u.(string); ok { + profile.URIs = append(profile.URIs, s) + } + } + } + } + if rh, ok := params["request_headers"]; ok { + if mp, ok := rh.(map[string]interface{}); ok { + profile.RequestHeaders = make(map[string]string) + for k, v := range mp { + profile.RequestHeaders[k], _ = v.(string) + } + } + } + if rh, ok := params["response_headers"]; ok { + if mp, ok := rh.(map[string]interface{}); ok { + profile.ResponseHeaders = make(map[string]string) + for k, v := range mp { + profile.ResponseHeaders[k], _ = v.(string) + } + } + } + if err := m.DB().UpdateC2Profile(profile); err != nil { + return makeC2Result(nil, err) + } + return makeC2Result(map[string]interface{}{"profile": profile}, nil) + + case "delete": + err := m.DB().DeleteC2Profile(id) + return makeC2Result(map[string]interface{}{"deleted": err == nil}, err) + + default: + return makeC2Result(nil, fmt.Errorf("unknown action: %s", action)) + } + }) +} + +// ============================================================================ +// c2_file — 文件管理工具(新增) +// ============================================================================ + +func registerC2FileTool(s *mcp.Server, m *c2.Manager, l *zap.Logger) { + s.RegisterTool(mcp.Tool{ + Name: builtin.ToolC2File, + Description: `C2 文件管理。通过 action 参数选择操作: +- list: 列出会话的文件传输记录(需 session_id) +- get_result: 获取任务结果文件路径(截图等,需 task_id)`, + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "action": map[string]interface{}{"type": "string", "description": "操作: list/get_result", "enum": []string{"list", "get_result"}}, + "session_id": map[string]interface{}{"type": "string", "description": "会话 ID(list 需要)"}, + "task_id": map[string]interface{}{"type": "string", "description": "任务 ID(get_result 需要)"}, + }, + "required": []string{"action"}, + }, + }, func(ctx context.Context, params map[string]interface{}) (*mcp.ToolResult, error) { + action := getString(params, "action") + + switch action { + case "list": + sessionID := getString(params, "session_id") + if sessionID == "" { + return makeC2Result(nil, fmt.Errorf("session_id required")) + } + files, err := m.DB().ListC2FilesBySession(sessionID) + return makeC2Result(map[string]interface{}{"files": files, "count": len(files)}, err) + + case "get_result": + taskID := getString(params, "task_id") + task, err := m.DB().GetC2Task(taskID) + if err != nil { + return makeC2Result(nil, err) + } + if task == nil { + return makeC2Result(nil, fmt.Errorf("task not found")) + } + if task.ResultBlobPath == "" { + return makeC2Result(map[string]interface{}{"has_file": false, "task_id": taskID}, nil) + } + return makeC2Result(map[string]interface{}{ + "has_file": true, + "task_id": taskID, + "file_path": task.ResultBlobPath, + }, nil) + + default: + return makeC2Result(nil, fmt.Errorf("unknown action: %s", action)) + } + }) +} + +// ============================================================================ +// 工具函数 +// ============================================================================ + +func getString(params map[string]interface{}, key string) string { + if v, ok := params[key]; ok { + if s, ok := v.(string); ok { + return s + } + } + return "" +} + +func getFloat64(params map[string]interface{}, key string) float64 { + if v, ok := params[key]; ok { + switch n := v.(type) { + case float64: + return n + case int: + return float64(n) + case string: + if f, err := strconv.ParseFloat(n, 64); err == nil { + return f + } + } + } + return 0 +} diff --git a/internal/app/c2_tools_test.go b/internal/app/c2_tools_test.go new file mode 100644 index 00000000..b24d2660 --- /dev/null +++ b/internal/app/c2_tools_test.go @@ -0,0 +1,69 @@ +package app + +import ( + "context" + "path/filepath" + "testing" + + "cyberstrike-ai/internal/authctx" + "cyberstrike-ai/internal/c2" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" + + "go.uber.org/zap" +) + +func TestC2ListenerCreateInheritsConversationProject(t *testing.T) { + db, err := database.NewDB(filepath.Join(t.TempDir(), "c2-tools.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + user, err := db.CreateRBACUser("c2-agent", "C2 Agent", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + project, err := db.CreateProject(&database.Project{Name: "engagement"}) + if err != nil { + t.Fatal(err) + } + if err := db.AssignResourceToUser(user.ID, "project", project.ID); err != nil { + t.Fatal(err) + } + conversation, err := db.CreateConversation("project chat", database.ConversationCreateMeta{ProjectID: project.ID}) + if err != nil { + t.Fatal(err) + } + + principal := authctx.NewPrincipal(user.ID, user.Username, database.RBACScopeAssigned, map[string]bool{ + "c2:read": true, "c2:write": true, + }) + ctx := authctx.WithPrincipal(mcp.WithMCPConversationID(context.Background(), conversation.ID), principal) + server := mcp.NewServer(zap.NewNop()) + server.SetToolAuthorizer(mcpToolAuthorizer(db)) + registerC2Tools(server, c2.NewManager(db, zap.NewNop(), t.TempDir()), zap.NewNop(), 8080) + + result, _, err := server.CallTool(ctx, builtin.ToolC2Listener, map[string]interface{}{ + "action": "create", + "name": "tcp-reverse-2222", + "type": "tcp_reverse", + "bind_host": "0.0.0.0", + "bind_port": 2222, + }) + if err != nil || result == nil || result.IsError { + t.Fatalf("create listener result=%#v err=%v text=%q", result, err, toolResultText(result)) + } + + listeners, err := db.ListC2Listeners() + if err != nil { + t.Fatal(err) + } + if len(listeners) != 1 { + t.Fatalf("listener count=%d, want 1", len(listeners)) + } + if listeners[0].ProjectID != project.ID { + t.Fatalf("listener project_id=%q, want %q", listeners[0].ProjectID, project.ID) + } +} diff --git a/internal/app/cors_security_test.go b/internal/app/cors_security_test.go new file mode 100644 index 00000000..d58eebf4 --- /dev/null +++ b/internal/app/cors_security_test.go @@ -0,0 +1,105 @@ +package app + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +func TestCORSMiddlewareAllowsSameOriginAndRejectsForeignOrigin(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(corsMiddleware(nil)) + router.GET("/test", func(c *gin.Context) { c.Status(http.StatusNoContent) }) + + same := httptest.NewRequest(http.MethodGet, "http://app.example/test", nil) + same.Host = "app.example" + same.Header.Set("Origin", "http://app.example") + sameW := httptest.NewRecorder() + router.ServeHTTP(sameW, same) + if sameW.Code != http.StatusNoContent || sameW.Header().Get("Access-Control-Allow-Origin") != "http://app.example" { + t.Fatalf("same-origin response = %d, allow-origin=%q", sameW.Code, sameW.Header().Get("Access-Control-Allow-Origin")) + } + + foreign := httptest.NewRequest(http.MethodGet, "http://app.example/test", nil) + foreign.Host = "app.example" + foreign.Header.Set("Origin", "https://evil.example") + foreignW := httptest.NewRecorder() + router.ServeHTTP(foreignW, foreign) + if foreignW.Code != http.StatusForbidden { + t.Fatalf("foreign-origin response = %d, want %d", foreignW.Code, http.StatusForbidden) + } +} + +func TestCORSMiddlewareAllowsBrowserExtensionWithoutConfiguration(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(corsMiddleware(nil)) + router.POST("/api/auth/login", func(c *gin.Context) { c.Status(http.StatusNoContent) }) + + req := httptest.NewRequest(http.MethodOptions, "https://server.example/api/auth/login", nil) + req.Host = "server.example" + req.Header.Set("Origin", "chrome-extension://abcdefghijklmnopabcdefghijklmnop") + req.Header.Set("Access-Control-Request-Method", http.MethodPost) + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + if w.Code != http.StatusNoContent { + t.Fatalf("preflight response = %d, want %d", w.Code, http.StatusNoContent) + } + if got := w.Header().Get("Access-Control-Allow-Origin"); got != "chrome-extension://abcdefghijklmnopabcdefghijklmnop" { + t.Fatalf("allow-origin = %q", got) + } +} + +func TestCORSMiddlewareRejectsInvalidExtensionOrigins(t *testing.T) { + gin.SetMode(gin.TestMode) + for _, origin := range []string{ + "chrome-extension://too-short", + "chrome-extension://qrstuvwxyzabcdefqrstuvwxyzabcdef", + "chrome-extension://abcdefghijklmnopabcdefghijklmnop:8443", + "moz-extension://abcdefghijklmnopabcdefghijklmnop", + } { + t.Run(origin, func(t *testing.T) { + router := gin.New() + router.Use(corsMiddleware(nil)) + router.GET("/test", func(c *gin.Context) { c.Status(http.StatusNoContent) }) + + req := httptest.NewRequest(http.MethodGet, "https://server.example/test", nil) + req.Host = "server.example" + req.Header.Set("Origin", origin) + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + if w.Code != http.StatusForbidden { + t.Fatalf("response = %d, want %d", w.Code, http.StatusForbidden) + } + }) + } +} + +func TestCORSMiddlewareRejectsUnsafeConfiguredEntries(t *testing.T) { + gin.SetMode(gin.TestMode) + for _, configured := range []string{ + "*", + "null", + "https://trusted.example/extra", + "https://trusted.example?trusted=true", + } { + t.Run(configured, func(t *testing.T) { + router := gin.New() + router.Use(corsMiddleware([]string{configured})) + router.GET("/test", func(c *gin.Context) { c.Status(http.StatusNoContent) }) + + req := httptest.NewRequest(http.MethodGet, "https://server.example/test", nil) + req.Host = "server.example" + req.Header.Set("Origin", "https://trusted.example") + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + if w.Code != http.StatusForbidden { + t.Fatalf("response = %d, want %d", w.Code, http.StatusForbidden) + } + }) + } +} diff --git a/internal/app/main_server_http_redirect.go b/internal/app/main_server_http_redirect.go new file mode 100644 index 00000000..7c7b74d7 --- /dev/null +++ b/internal/app/main_server_http_redirect.go @@ -0,0 +1,213 @@ +package app + +import ( + "bufio" + "context" + "crypto/tls" + "errors" + "fmt" + "net" + "net/http" + "strconv" + "sync" + "time" + + "go.uber.org/zap" +) + +// peekedConn 在已预读首字节后仍将连接交给 net/http 或 crypto/tls。 +type peekedConn struct { + net.Conn + r *bufio.Reader +} + +func (c *peekedConn) Read(p []byte) (int, error) { + return c.r.Read(p) +} + +// oneConnListener 供 http.Server.Serve 处理单条 TCP 连接(含 keep-alive)。 +type oneConnListener struct { + conn net.Conn + addr net.Addr + once sync.Once +} + +func (l *oneConnListener) Accept() (net.Conn, error) { + var c net.Conn + l.once.Do(func() { + c = l.conn + l.conn = nil + }) + if c == nil { + return nil, net.ErrClosed + } + return c, nil +} + +func (l *oneConnListener) Close() error { return nil } +func (l *oneConnListener) Addr() net.Addr { return l.addr } + +// httpServerForTLSConn 从已有 Server 复制可服务字段,用于已握手 TLS 连接上的 HTTP 服务。 +// 不能复制整个 http.Server(内含 atomic/noCopy 字段)。 +func httpServerForTLSConn(src *http.Server) *http.Server { + return &http.Server{ + Handler: src.Handler, + DisableGeneralOptionsHandler: src.DisableGeneralOptionsHandler, + ReadTimeout: src.ReadTimeout, + ReadHeaderTimeout: src.ReadHeaderTimeout, + WriteTimeout: src.WriteTimeout, + IdleTimeout: src.IdleTimeout, + MaxHeaderBytes: src.MaxHeaderBytes, + ConnState: src.ConnState, + ErrorLog: src.ErrorLog, + BaseContext: src.BaseContext, + ConnContext: src.ConnContext, + } +} + +func isTLSHandshakeRecord(b byte) bool { + return b == 0x16 +} + +func newHTTPToHTTPSRedirectHandler(httpsPort int) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + host := r.Host + if h, _, err := net.SplitHostPort(host); err == nil { + host = h + } + var target string + if httpsPort == 443 { + target = fmt.Sprintf("https://%s%s", host, r.URL.RequestURI()) + } else { + target = fmt.Sprintf("https://%s:%d%s", host, httpsPort, r.URL.RequestURI()) + } + http.Redirect(w, r, target, http.StatusPermanentRedirect) + }) +} + +func portFromListenAddr(addr string) int { + _, portStr, err := net.SplitHostPort(addr) + if err != nil { + return 443 + } + p, err := strconv.Atoi(portStr) + if err != nil || p <= 0 { + return 443 + } + return p +} + +func ensureMainTLSConfigCerts(mode mainTLSMode, tlsConf *tls.Config, certFile, keyFile string) (*tls.Config, error) { + if mode != mainTLSFromFiles { + return tlsConf, nil + } + if tlsConf == nil { + tlsConf = &tls.Config{MinVersion: tls.VersionTLS12} + } + if len(tlsConf.Certificates) > 0 { + return tlsConf, nil + } + cert, err := tls.LoadX509KeyPair(certFile, keyFile) + if err != nil { + return nil, err + } + tlsConf.Certificates = []tls.Certificate{cert} + return tlsConf, nil +} + +type mainServerMux struct { + ln net.Listener + httpsSrv *http.Server + redirectSrv *http.Server + logger *zap.Logger +} + +func newMainServerMux(ln net.Listener, httpsSrv *http.Server, httpsPort int, logger *zap.Logger) *mainServerMux { + return &mainServerMux{ + ln: ln, + httpsSrv: httpsSrv, + redirectSrv: &http.Server{Handler: newHTTPToHTTPSRedirectHandler(httpsPort), ReadHeaderTimeout: 10 * time.Second}, + logger: logger, + } +} + +func (m *mainServerMux) Serve() error { + for { + conn, err := m.ln.Accept() + if err != nil { + if errors.Is(err, net.ErrClosed) { + return http.ErrServerClosed + } + return err + } + go m.handleConn(conn) + } +} + +func (m *mainServerMux) handleConn(raw net.Conn) { + if err := raw.SetReadDeadline(time.Now().Add(10 * time.Second)); err != nil { + _ = raw.Close() + return + } + br := bufio.NewReader(raw) + b, err := br.Peek(1) + if err != nil { + _ = raw.Close() + return + } + _ = raw.SetReadDeadline(time.Time{}) + + pc := &peekedConn{Conn: raw, r: br} + ocl := &oneConnListener{conn: pc, addr: raw.LocalAddr()} + + if isTLSHandshakeRecord(b[0]) { + m.serveHTTPS(pc, raw.LocalAddr()) + return + } + if err := m.redirectSrv.Serve(ocl); err != nil && !errors.Is(err, net.ErrClosed) && !errors.Is(err, http.ErrServerClosed) { + m.logger.Debug("HTTP 重定向连接处理结束", zap.Error(err)) + } +} + +// serveHTTPS 在已嗅探为 TLS 的连接上完成握手,再按 ALPN 走 HTTP/2 或 HTTP/1.1。 +// 不能对同一 http.Server 并发调用 Serve(TLSConfig!=nil),否则握手/ALPN 会异常(浏览器 ERR_SSL_PROTOCOL_ERROR)。 +func (m *mainServerMux) serveHTTPS(pc *peekedConn, localAddr net.Addr) { + tlsConn := tls.Server(pc, m.httpsSrv.TLSConfig) + handCtx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + if err := tlsConn.HandshakeContext(handCtx); err != nil { + m.logger.Debug("TLS 握手失败", zap.Error(err)) + _ = pc.Close() + return + } + + srv := m.httpsSrv + if srv.TLSNextProto != nil { + proto := tlsConn.ConnectionState().NegotiatedProtocol + if fn := srv.TLSNextProto[proto]; fn != nil { + fn(srv, tlsConn, srv.Handler) + return + } + } + + plain := httpServerForTLSConn(srv) + ocl := &oneConnListener{conn: tlsConn, addr: localAddr} + if err := plain.Serve(ocl); err != nil && !errors.Is(err, net.ErrClosed) && !errors.Is(err, http.ErrServerClosed) { + m.logger.Debug("HTTPS 连接处理结束", zap.Error(err)) + } +} + +func (m *mainServerMux) Shutdown(ctx context.Context) error { + _ = m.ln.Close() + var err1, err2 error + if m.httpsSrv != nil { + err1 = m.httpsSrv.Shutdown(ctx) + } + if m.redirectSrv != nil { + err2 = m.redirectSrv.Shutdown(ctx) + } + if err1 != nil { + return err1 + } + return err2 +} diff --git a/internal/app/main_server_http_redirect_test.go b/internal/app/main_server_http_redirect_test.go new file mode 100644 index 00000000..99037f29 --- /dev/null +++ b/internal/app/main_server_http_redirect_test.go @@ -0,0 +1,150 @@ +package app + +import ( + "crypto/tls" + "io" + "net" + "net/http" + "net/http/httptest" + "strconv" + "testing" + + "cyberstrike-ai/internal/config" + + "golang.org/x/net/http2" +) + +func TestNewHTTPToHTTPSRedirectHandler(t *testing.T) { + t.Parallel() + tests := []struct { + name string + httpsPort int + host string + uri string + wantTarget string + }{ + { + name: "non standard port", + httpsPort: 8080, + host: "127.0.0.1:8080", + uri: "/login?next=/", + wantTarget: "https://127.0.0.1:8080/login?next=/", + }, + { + name: "standard port", + httpsPort: 443, + host: "example.com:80", + uri: "/", + wantTarget: "https://example.com/", + }, + } + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + h := newHTTPToHTTPSRedirectHandler(tt.httpsPort) + req := httptest.NewRequest(http.MethodGet, "http://"+tt.host+tt.uri, nil) + req.Host = tt.host + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + if rec.Code != http.StatusPermanentRedirect { + t.Fatalf("status = %d, want %d", rec.Code, http.StatusPermanentRedirect) + } + if got := rec.Header().Get("Location"); got != tt.wantTarget { + t.Fatalf("Location = %q, want %q", got, tt.wantTarget) + } + }) + } +} + +func TestIsTLSHandshakeRecord(t *testing.T) { + t.Parallel() + if !isTLSHandshakeRecord(0x16) { + t.Fatal("expected TLS handshake record") + } + if isTLSHandshakeRecord('G') { + t.Fatal("GET should not be TLS") + } +} + +func TestServerHTTPRedirectEnabled(t *testing.T) { + t.Parallel() + disabled := false + enabled := true + if config.ServerHTTPRedirectEnabled(nil) { + t.Fatal("nil config should disable redirect") + } + if !config.ServerHTTPRedirectEnabled(&config.ServerConfig{TLSEnabled: true}) { + t.Fatal("HTTPS without explicit flag should enable redirect") + } + if config.ServerHTTPRedirectEnabled(&config.ServerConfig{TLSEnabled: true, TLSHTTPRedirect: &disabled}) { + t.Fatal("explicit false should disable redirect") + } + if !config.ServerHTTPRedirectEnabled(&config.ServerConfig{TLSEnabled: true, TLSHTTPRedirect: &enabled}) { + t.Fatal("explicit true should enable redirect") + } + if config.ServerHTTPRedirectEnabled(&config.ServerConfig{}) { + t.Fatal("plain HTTP should not redirect") + } +} + +func TestMainServerMuxHTTPRedirectAndHTTPS(t *testing.T) { + cert, err := generateMainServerSelfSignedCert() + if err != nil { + t.Fatalf("generate cert: %v", err) + } + handler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, "ok") + }) + srv := &http.Server{Handler: handler, TLSConfig: &tls.Config{ + MinVersion: tls.VersionTLS12, + Certificates: []tls.Certificate{cert}, + }} + if err := http2.ConfigureServer(srv, &http2.Server{}); err != nil { + t.Fatalf("configure http2: %v", err) + } + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + defer ln.Close() + + mux := newMainServerMux(ln, srv, portFromListenAddr(ln.Addr().String()), nil) + go func() { _ = mux.Serve() }() + + client := &http.Client{ + Transport: &http.Transport{ + TLSClientConfig: &tls.Config{InsecureSkipVerify: true, MinVersion: tls.VersionTLS12}, + }, + CheckRedirect: func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + }, + } + addr := ln.Addr().String() + + httpResp, err := client.Get("http://" + addr + "/") + if err != nil { + t.Fatalf("http get: %v", err) + } + _ = httpResp.Body.Close() + if httpResp.StatusCode != http.StatusPermanentRedirect { + t.Fatalf("http status = %d, want %d", httpResp.StatusCode, http.StatusPermanentRedirect) + } + if got := httpResp.Header.Get("Location"); got != "https://127.0.0.1:"+strconv.Itoa(portFromListenAddr(addr))+"/" { + t.Fatalf("Location = %q", got) + } + + httpsResp, err := client.Get("https://" + addr + "/") + if err != nil { + t.Fatalf("https get: %v", err) + } + defer httpsResp.Body.Close() + if httpsResp.StatusCode != http.StatusOK { + t.Fatalf("https status = %d, want %d", httpsResp.StatusCode, http.StatusOK) + } + body, _ := io.ReadAll(httpsResp.Body) + if string(body) != "ok" { + t.Fatalf("body = %q, want ok", body) + } +} diff --git a/internal/app/main_server_tls.go b/internal/app/main_server_tls.go new file mode 100644 index 00000000..19b546d6 --- /dev/null +++ b/internal/app/main_server_tls.go @@ -0,0 +1,86 @@ +package app + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "fmt" + "math/big" + "net" + "strings" + "time" + + "cyberstrike-ai/internal/config" +) + +// mainTLSMode 主 Web 服务 TLS 启动方式。 +type mainTLSMode int + +const ( + mainTLSOff mainTLSMode = iota + mainTLSFromFiles + mainTLSInMemorySelfSigned +) + +// prepareMainServerTLS 根据 server 配置决定主站是否启用 HTTPS(及 HTTP/2 协商)。 +// fromFiles:使用 tls_cert_path + tls_key_path,由 http.Server.ListenAndServeTLS 加载 PEM。 +// inMemory:tls_auto_self_sign 生成的自签证书,仅用于本地/测试。 +func prepareMainServerTLS(cfg *config.ServerConfig) (mode mainTLSMode, tlsConf *tls.Config, certFile, keyFile string, err error) { + if cfg == nil || !config.MainWebUIUsesHTTPS(cfg) { + return mainTLSOff, nil, "", "", nil + } + certFile = strings.TrimSpace(cfg.TLSCertPath) + keyFile = strings.TrimSpace(cfg.TLSKeyPath) + if certFile != "" && keyFile != "" { + // 证书由 ListenAndServeTLS 从文件加载;此处仅提供最小 TLS 配置供 http2.ConfigureServer 合并 ALPN。 + return mainTLSFromFiles, &tls.Config{MinVersion: tls.VersionTLS12}, certFile, keyFile, nil + } + if cfg.TLSAutoSelfSign { + cert, genErr := generateMainServerSelfSignedCert() + if genErr != nil { + return mainTLSOff, nil, "", "", fmt.Errorf("生成自签 TLS 证书: %w", genErr) + } + tlsConf = &tls.Config{ + MinVersion: tls.VersionTLS12, + Certificates: []tls.Certificate{cert}, + } + return mainTLSInMemorySelfSigned, tlsConf, "", "", nil + } + return mainTLSOff, nil, "", "", fmt.Errorf("server: 已启用 TLS(tls_enabled / tls_auto_self_sign / 证书路径),请设置 tls_cert_path 与 tls_key_path,或将 tls_auto_self_sign 设为 true(仅测试环境)") +} + +func generateMainServerSelfSignedCert() (tls.Certificate, error) { + priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + return tls.Certificate{}, err + } + serial, err := rand.Int(rand.Reader, big.NewInt(1<<62)) + if err != nil { + return tls.Certificate{}, err + } + tmpl := &x509.Certificate{ + SerialNumber: serial, + Subject: pkix.Name{CommonName: "CyberStrikeAI"}, + NotBefore: time.Now().Add(-1 * time.Hour), + NotAfter: time.Now().Add(365 * 24 * time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + IPAddresses: []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("::1")}, + DNSNames: []string{"localhost"}, + } + der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &priv.PublicKey, priv) + if err != nil { + return tls.Certificate{}, err + } + keyDER, err := x509.MarshalECPrivateKey(priv) + if err != nil { + return tls.Certificate{}, err + } + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}) + return tls.X509KeyPair(certPEM, keyPEM) +} diff --git a/internal/app/mcp_authorization.go b/internal/app/mcp_authorization.go new file mode 100644 index 00000000..6526de7d --- /dev/null +++ b/internal/app/mcp_authorization.go @@ -0,0 +1,407 @@ +package app + +import ( + "context" + "fmt" + "strings" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/authctx" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" +) + +func mcpToolAuthorizer(db *database.DB) func(context.Context, string, map[string]interface{}) error { + return func(ctx context.Context, toolName string, args map[string]interface{}) error { + principal, ok := authctx.PrincipalFromContext(ctx) + if !ok { + return fmt.Errorf("missing authenticated principal") + } + require := func(permission string) error { + if !principal.HasPermission(permission) { + return fmt.Errorf("missing permission %s", permission) + } + return nil + } + resource := func(permission, resourceType, argument string) error { + if err := require(permission); err != nil { + return err + } + id := mcpAuthorizationString(args, argument) + if id == "" || db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), resourceType, id) { + return fmt.Errorf("no access to %s %s", resourceType, id) + } + if err := authorizeMCPProjectResourceBoundary(ctx, db, resourceType, id); err != nil { + return err + } + return nil + } + toolExecutionResource := func(permission string) error { + if err := require(permission); err != nil { + return err + } + id := mcpAuthorizationString(args, "execution_id") + if id == "" || db == nil || !db.UserCanAccessToolExecution(principal.UserID, principal.ScopeFor(permission), id) { + return fmt.Errorf("no access to tool execution %s", id) + } + return nil + } + + switch toolName { + case builtin.ToolWebshellExec, builtin.ToolWebshellFileWrite: + return resource("webshell:write", "webshell", "connection_id") + case builtin.ToolWebshellFileList, builtin.ToolWebshellFileRead: + return resource("webshell:read", "webshell", "connection_id") + case builtin.ToolManageWebshellList: + return require("webshell:read") + case builtin.ToolManageWebshellAdd: + return require("webshell:write") + case builtin.ToolManageWebshellUpdate, builtin.ToolManageWebshellTest: + return resource("webshell:write", "webshell", "connection_id") + case builtin.ToolManageWebshellDelete: + return resource("webshell:delete", "webshell", "connection_id") + case builtin.ToolRecordVulnerability: + if err := require("vulnerability:write"); err != nil { + return err + } + conversationID := mcpAuthorizationString(args, "conversation_id") + if conversationID == "" { + conversationID = mcpAuthorizationConversationID(ctx) + } + if conversationID == "" || db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor("vulnerability:write"), "conversation", conversationID) { + return fmt.Errorf("no access to conversation %s", conversationID) + } + return nil + case builtin.ToolListVulnerabilities: + if err := require("vulnerability:read"); err != nil { + return err + } + conversationID := mcpAuthorizationConversationID(ctx) + if conversationID == "" || db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor("vulnerability:read"), "conversation", conversationID) { + return fmt.Errorf("no access to conversation %s", conversationID) + } + return nil + case builtin.ToolGetVulnerability: + return resource("vulnerability:read", "vulnerability", "id") + case builtin.ToolQueryAssets: + return require("asset:read") + case builtin.ToolGetAsset: + return resource("asset:read", "asset", "id") + case builtin.ToolCreateAsset: + if err := require("asset:write"); err != nil { + return err + } + if projectID := mcpAuthorizationString(args, "project_id"); projectID != "" && (db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor("asset:write"), "project", projectID)) { + return fmt.Errorf("no access to project %s", projectID) + } + return nil + case builtin.ToolUpdateAsset, builtin.ToolCompleteAssetScan: + if err := resource("asset:write", "asset", "id"); err != nil { + return err + } + if toolName == builtin.ToolCompleteAssetScan { + conversationID := mcpAuthorizationConversationID(ctx) + if conversationID == "" || db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor("asset:write"), "conversation", conversationID) { + return fmt.Errorf("no access to conversation %s", conversationID) + } + return nil + } + if projectID := mcpAuthorizationString(args, "project_id"); projectID != "" && (db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor("asset:write"), "project", projectID)) { + return fmt.Errorf("no access to project %s", projectID) + } + return nil + case builtin.ToolDeleteAsset: + return resource("asset:delete", "asset", "id") + case builtin.ToolUpsertProjectFact, builtin.ToolDeprecateProjectFact, builtin.ToolRestoreProjectFact: + return authorizeProjectTool(ctx, principal, db, "project:write") + case builtin.ToolGetProjectFact, builtin.ToolListProjectFacts, builtin.ToolSearchProjectFacts: + return authorizeProjectTool(ctx, principal, db, "project:read") + case builtin.ToolListKnowledgeRiskTypes, builtin.ToolSearchKnowledgeBase: + return require("knowledge:read") + case builtin.ToolAnalyzeImage: + return require("agent:execute") + case builtin.ToolGetToolExecution, builtin.ToolWaitToolExecution: + return toolExecutionResource("monitor:read") + case builtin.ToolCancelToolExecution: + return toolExecutionResource("monitor:write") + case builtin.ToolBatchTaskList: + return require("tasks:read") + case builtin.ToolBatchTaskGet: + return resource("tasks:read", "batch_task", "queue_id") + case builtin.ToolBatchTaskCreate: + if err := require("tasks:write"); err != nil { + return err + } + if projectID := mcpAuthorizationString(args, "project_id"); projectID != "" && (db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor("tasks:write"), "project", projectID)) { + return fmt.Errorf("no access to project %s", projectID) + } + return nil + case builtin.ToolBatchTaskDelete, builtin.ToolBatchTaskRemove: + return resource("tasks:delete", "batch_task", "queue_id") + case builtin.ToolBatchTaskStart, builtin.ToolBatchTaskRerun, builtin.ToolBatchTaskPause, + builtin.ToolBatchTaskUpdateMetadata, builtin.ToolBatchTaskUpdateSchedule, + builtin.ToolBatchTaskScheduleEnabled, builtin.ToolBatchTaskAdd, builtin.ToolBatchTaskUpdate: + return resource("tasks:write", "batch_task", "queue_id") + case builtin.ToolC2Listener: + return authorizeC2Action(ctx, principal, db, args, "c2_listener", "listener_id") + case builtin.ToolC2Session, builtin.ToolC2Task, builtin.ToolC2File: + if toolName == builtin.ToolC2File && mcpAuthorizationString(args, "action") == "get_result" { + return authorizeC2Action(ctx, principal, db, args, "c2_task", "task_id") + } + return authorizeC2Action(ctx, principal, db, args, "c2_session", "session_id") + case builtin.ToolC2TaskManage: + return authorizeC2Action(ctx, principal, db, args, "c2_task", "task_id") + case builtin.ToolC2Payload: + return resource("c2:write", "c2_listener", "listener_id") + case builtin.ToolC2Event: + if id := mcpAuthorizationString(args, "session_id"); id != "" { + return resource("c2:read", "c2_session", "session_id") + } + if id := mcpAuthorizationString(args, "task_id"); id != "" { + return resource("c2:read", "c2_task", "task_id") + } + if filter := mcpEffectiveProjectFilter(ctx, db); filter != "" { + return require("c2:read") + } + if principal.ScopeFor("c2:read") != database.RBACScopeAll { + return fmt.Errorf("unfiltered C2 event list requires global scope") + } + return require("c2:read") + case builtin.ToolC2Profile: + // Profiles are process-global and do not yet have an owner. Writes are + // therefore reserved for global scope; reads require c2:read. + if mcpAuthorizationString(args, "action") == "list" || mcpAuthorizationString(args, "action") == "get" { + return require("c2:read") + } + permission := "c2:write" + if mcpAuthorizationString(args, "action") == "delete" { + permission = "c2:delete" + } + if principal.ScopeFor(permission) != database.RBACScopeAll { + return fmt.Errorf("C2 profile mutation requires global scope") + } + if mcpAuthorizationString(args, "action") == "delete" { + return require("c2:delete") + } + return require("c2:write") + default: + if builtin.IsBuiltinTool(toolName) { + return fmt.Errorf("no authorization policy registered for builtin tool %s", toolName) + } + if principal.HasPermission("agent:local-execute") { + return nil + } + return fmt.Errorf("missing agent:local-execute") + } + } +} + +func externalMCPToolAuthorizer() func(context.Context, string, map[string]interface{}) error { + return func(ctx context.Context, toolName string, _ map[string]interface{}) error { + principal, ok := authctx.PrincipalFromContext(ctx) + if !ok { + return fmt.Errorf("missing authenticated principal") + } + if !principal.HasPermission("mcp:external:execute") { + return fmt.Errorf("missing permission mcp:external:execute") + } + if principal.ScopeFor("mcp:external:execute") != database.RBACScopeAll { + return fmt.Errorf("external MCP invocation requires global scope") + } + if strings.TrimSpace(toolName) == "" { + return fmt.Errorf("missing external tool name") + } + return nil + } +} + +func authorizeC2Action(ctx context.Context, principal authctx.Principal, db *database.DB, args map[string]interface{}, resourceType, argument string) error { + action := mcpAuthorizationString(args, "action") + permission := "c2:write" + if action == "list" || action == "get" || action == "get_result" || action == "wait" { + permission = "c2:read" + } else if action == "delete" || action == "delete_batch" { + permission = "c2:delete" + } + if !principal.HasPermission(permission) { + return fmt.Errorf("missing permission %s", permission) + } + id := mcpAuthorizationString(args, argument) + if action == "delete_batch" { + ids := mcpAuthorizationStrings(args, argument+"s") + if len(ids) == 0 { + return fmt.Errorf("missing resource identifiers %ss", argument) + } + for _, candidate := range ids { + if db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), resourceType, candidate) { + return fmt.Errorf("no access to %s %s", resourceType, candidate) + } + if err := authorizeMCPProjectResourceBoundary(ctx, db, resourceType, candidate); err != nil { + return err + } + } + return nil + } + if id == "" { + if action == "create" { + projectID := mcpAuthorizationString(args, "project_id") + if projectID == "" { + projectID = mcpEffectiveProjectFilter(ctx, db) + if projectID == database.ProjectFilterUnbound { + projectID = "" + } + } + if projectID != "" && (db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), "project", projectID)) { + return fmt.Errorf("no access to project %s", projectID) + } + return nil + } + if action == "list" { + return nil + } + return fmt.Errorf("missing resource identifier %s", argument) + } + if db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), resourceType, id) { + return fmt.Errorf("no access to %s %s", resourceType, id) + } + if err := authorizeMCPProjectResourceBoundary(ctx, db, resourceType, id); err != nil { + return err + } + return nil +} + +func authorizeMCPProjectResourceBoundary(ctx context.Context, db *database.DB, resourceType, resourceID string) error { + filter := mcpEffectiveProjectFilter(ctx, db) + if filter == "" || db == nil { + return nil + } + projectID, ok, err := mcpResourceProjectID(db, resourceType, resourceID) + if err != nil { + return err + } + if !ok { + return nil + } + if filter == database.ProjectFilterUnbound { + if projectID != "" { + return fmt.Errorf("resource %s %s belongs to project %s, current conversation is unbound", resourceType, resourceID, projectID) + } + return nil + } + if projectID != filter { + if projectID == "" { + return fmt.Errorf("resource %s %s is unbound, current conversation project is %s", resourceType, resourceID, filter) + } + return fmt.Errorf("resource %s %s belongs to project %s, current conversation project is %s", resourceType, resourceID, projectID, filter) + } + return nil +} + +func mcpResourceProjectID(db *database.DB, resourceType, resourceID string) (string, bool, error) { + switch resourceType { + case "webshell": + conn, err := db.GetWebshellConnection(resourceID) + if err != nil { + return "", true, err + } + if conn == nil { + return "", true, fmt.Errorf("webshell not found") + } + return strings.TrimSpace(conn.ProjectID), true, nil + case "c2_listener": + listener, err := db.GetC2Listener(resourceID) + if err != nil { + return "", true, err + } + if listener == nil { + return "", true, fmt.Errorf("listener not found") + } + return strings.TrimSpace(listener.ProjectID), true, nil + case "c2_session": + session, err := db.GetC2Session(resourceID) + if err != nil { + return "", true, err + } + if session == nil { + return "", true, fmt.Errorf("session not found") + } + return mcpResourceProjectID(db, "c2_listener", session.ListenerID) + case "c2_task": + task, err := db.GetC2Task(resourceID) + if err != nil { + return "", true, err + } + if task == nil { + return "", true, fmt.Errorf("task not found") + } + return mcpResourceProjectIDFromC2Session(db, task.SessionID) + default: + return "", false, nil + } +} + +func mcpResourceProjectIDFromC2Session(db *database.DB, sessionID string) (string, bool, error) { + session, err := db.GetC2Session(sessionID) + if err != nil { + return "", true, err + } + if session == nil { + return "", true, fmt.Errorf("session not found") + } + return mcpResourceProjectID(db, "c2_listener", session.ListenerID) +} + +func mcpAuthorizationStrings(args map[string]interface{}, key string) []string { + values := []string{} + switch raw := args[key].(type) { + case []string: + for _, value := range raw { + if value = strings.TrimSpace(value); value != "" { + values = append(values, value) + } + } + case []interface{}: + for _, item := range raw { + if value, ok := item.(string); ok { + if value = strings.TrimSpace(value); value != "" { + values = append(values, value) + } + } + } + } + return values +} + +func authorizeProjectTool(ctx context.Context, principal authctx.Principal, db *database.DB, permission string) error { + if !principal.HasPermission(permission) { + return fmt.Errorf("missing permission %s", permission) + } + conversationID := mcpAuthorizationConversationID(ctx) + if conversationID == "" || db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), "conversation", conversationID) { + return fmt.Errorf("no access to conversation %s", conversationID) + } + projectID, err := db.GetConversationProjectID(conversationID) + if err != nil { + return fmt.Errorf("no access to project: %w", err) + } + if strings.TrimSpace(projectID) == "" { + return fmt.Errorf("当前对话未绑定项目,无法使用项目黑板工具,请先在对话中选择项目或创建带项目的对话") + } + if !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), "project", projectID) { + return fmt.Errorf("no access to project %s", projectID) + } + return nil +} + +func mcpAuthorizationConversationID(ctx context.Context) string { + if id := strings.TrimSpace(agent.ConversationIDFromContext(ctx)); id != "" { + return id + } + return strings.TrimSpace(mcp.MCPConversationIDFromContext(ctx)) +} + +func mcpAuthorizationString(args map[string]interface{}, key string) string { + value, _ := args[key].(string) + return strings.TrimSpace(value) +} diff --git a/internal/app/mcp_authorization_test.go b/internal/app/mcp_authorization_test.go new file mode 100644 index 00000000..e4bfdc44 --- /dev/null +++ b/internal/app/mcp_authorization_test.go @@ -0,0 +1,261 @@ +package app + +import ( + "context" + "path/filepath" + "strings" + "testing" + "time" + + "cyberstrike-ai/internal/authctx" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" + "cyberstrike-ai/internal/security" + + "go.uber.org/zap" +) + +func TestMCPToolAuthorizerEnforcesPermissionAndResource(t *testing.T) { + db, err := database.NewDB(filepath.Join(t.TempDir(), "mcp-authz.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + user, err := db.CreateRBACUser("mcp-user", "MCP User", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + for _, id := range []string{"ws_allowed", "ws_hidden"} { + if err := db.CreateWebshellConnection(&database.WebShellConnection{ID: id, URL: "http://127.0.0.1/" + id, Type: "php", Method: "post", CmdParam: "cmd", CreatedAt: time.Now()}); err != nil { + t.Fatal(err) + } + } + if err := db.AssignResourceToUser(user.ID, "webshell", "ws_allowed"); err != nil { + t.Fatal(err) + } + + principal := authctx.NewPrincipal(user.ID, user.Username, database.RBACScopeAssigned, map[string]bool{"mcp:write": true, "webshell:write": true}) + ctx := authctx.WithPrincipal(context.Background(), principal) + authorize := mcpToolAuthorizer(db) + if err := authorize(ctx, builtin.ToolWebshellExec, map[string]interface{}{"connection_id": "ws_allowed"}); err != nil { + t.Fatalf("allowed resource denied: %v", err) + } + if err := authorize(ctx, builtin.ToolWebshellExec, map[string]interface{}{"connection_id": "ws_hidden"}); err == nil { + t.Fatal("foreign webshell resource was allowed") + } + if err := authorize(ctx, builtin.ToolManageWebshellDelete, map[string]interface{}{"connection_id": "ws_allowed"}); err == nil { + t.Fatal("delete without webshell:delete was allowed") + } +} + +func TestMCPToolAuthorizerEnforcesConversationProjectBoundary(t *testing.T) { + db, err := database.NewDB(filepath.Join(t.TempDir(), "mcp-project-boundary.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + user, err := db.CreateRBACUser("boundary-user", "Boundary User", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + project, err := db.CreateProject(&database.Project{Name: "Project 123"}) + if err != nil { + t.Fatal(err) + } + projectConv, err := db.CreateConversation("project conversation", database.ConversationCreateMeta{ProjectID: project.ID}) + if err != nil { + t.Fatal(err) + } + unboundConv, err := db.CreateConversation("unbound conversation", database.ConversationCreateMeta{}) + if err != nil { + t.Fatal(err) + } + + wsProject := database.WebShellConnection{ID: "ws_project", ProjectID: project.ID, URL: "http://127.0.0.1/project.php", Type: "php", Method: "post", CreatedAt: time.Now()} + wsUnbound := database.WebShellConnection{ID: "ws_unbound", URL: "http://127.0.0.1/unbound.php", Type: "php", Method: "post", CreatedAt: time.Now()} + if err := db.CreateWebshellConnection(&wsProject); err != nil { + t.Fatal(err) + } + if err := db.CreateWebshellConnection(&wsUnbound); err != nil { + t.Fatal(err) + } + for _, id := range []string{wsProject.ID, wsUnbound.ID} { + if err := db.AssignResourceToUser(user.ID, "webshell", id); err != nil { + t.Fatal(err) + } + } + + now := time.Now() + listener := &database.C2Listener{ID: "l_project", ProjectID: project.ID, Name: "project listener", Type: "tcp_reverse", BindHost: "127.0.0.1", BindPort: 5555, OwnerUserID: user.ID, CreatedAt: now} + if err := db.CreateC2Listener(listener); err != nil { + t.Fatal(err) + } + if err := db.AssignResourceToUser(user.ID, "c2_listener", listener.ID); err != nil { + t.Fatal(err) + } + session := &database.C2Session{ID: "s_project", ListenerID: listener.ID, ImplantUUID: "implant-project", Status: "active", FirstSeenAt: now, LastCheckIn: now} + if err := db.UpsertC2Session(session); err != nil { + t.Fatal(err) + } + + principal := authctx.NewPrincipal(user.ID, user.Username, database.RBACScopeAssigned, map[string]bool{ + "webshell:read": true, "webshell:write": true, + "c2:read": true, "c2:write": true, + }) + authorize := mcpToolAuthorizer(db) + unboundCtx := authctx.WithPrincipal(mcp.WithMCPConversationID(context.Background(), unboundConv.ID), principal) + projectCtx := authctx.WithPrincipal(mcp.WithMCPProjectID(mcp.WithMCPConversationID(context.Background(), projectConv.ID), project.ID), principal) + projectCtxFromConversationOnly := authctx.WithPrincipal(mcp.WithMCPConversationID(context.Background(), projectConv.ID), principal) + + if err := authorize(unboundCtx, builtin.ToolWebshellExec, map[string]interface{}{"connection_id": wsProject.ID}); err == nil { + t.Fatal("unbound conversation was allowed to use project-bound webshell") + } + if err := authorize(unboundCtx, builtin.ToolWebshellExec, map[string]interface{}{"connection_id": wsUnbound.ID}); err != nil { + t.Fatalf("unbound webshell denied in unbound conversation: %v", err) + } + if err := authorize(projectCtx, builtin.ToolWebshellExec, map[string]interface{}{"connection_id": wsProject.ID}); err != nil { + t.Fatalf("project webshell denied in project conversation: %v", err) + } + if err := authorize(projectCtxFromConversationOnly, builtin.ToolWebshellExec, map[string]interface{}{"connection_id": wsProject.ID}); err != nil { + t.Fatalf("project webshell denied when only conversation id is present: %v", err) + } + if err := authorize(projectCtx, builtin.ToolWebshellExec, map[string]interface{}{"connection_id": wsUnbound.ID}); err == nil { + t.Fatal("project conversation was allowed to use unbound webshell by id") + } + if err := authorize(unboundCtx, builtin.ToolC2Session, map[string]interface{}{"action": "get", "session_id": session.ID}); err == nil { + t.Fatal("unbound conversation was allowed to use project-bound c2 session") + } + if err := authorize(projectCtx, builtin.ToolC2Session, map[string]interface{}{"action": "get", "session_id": session.ID}); err != nil { + t.Fatalf("project c2 session denied in project conversation: %v", err) + } + if err := authorize(projectCtxFromConversationOnly, builtin.ToolC2Session, map[string]interface{}{"action": "get", "session_id": session.ID}); err != nil { + t.Fatalf("project c2 session denied when only conversation id is present: %v", err) + } +} + +func TestEveryBuiltinMCPToolHasExplicitAuthorizationPolicy(t *testing.T) { + db, err := database.NewDB(filepath.Join(t.TempDir(), "mcp-policy-inventory.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + permissions := map[string]bool{} + for permission := range security.PermissionCatalog { + permissions[permission] = true + } + ctx := authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("admin", "admin", database.RBACScopeAll, permissions)) + authorize := mcpToolAuthorizer(db) + args := map[string]interface{}{ + "action": "get", "connection_id": "x", "queue_id": "x", "listener_id": "x", + "session_id": "x", "task_id": "x", "id": "x", "conversation_id": "x", "execution_id": "x", + } + for _, toolName := range builtin.GetAllBuiltinTools() { + err := authorize(ctx, toolName, args) + if err != nil && strings.Contains(err.Error(), "no authorization policy registered") { + t.Errorf("builtin tool %s has no explicit policy", toolName) + } + } +} + +func TestMCPExecutionControlAuthorizationUsesExecutionScope(t *testing.T) { + db, err := database.NewDB(filepath.Join(t.TempDir(), "mcp-exec-authz.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + user, err := db.CreateRBACUser("exec-user", "Exec User", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + if err := db.SaveToolExecution(&mcp.ToolExecution{ + ID: "exec-owned", + ToolName: "lab::slow", + Status: "running", + StartTime: time.Now(), + OwnerUserID: user.ID, + }); err != nil { + t.Fatal(err) + } + if err := db.SaveToolExecution(&mcp.ToolExecution{ + ID: "exec-hidden", + ToolName: "lab::slow", + Status: "running", + StartTime: time.Now(), + }); err != nil { + t.Fatal(err) + } + + principal := authctx.NewPrincipal(user.ID, user.Username, database.RBACScopeAssigned, map[string]bool{"monitor:read": true, "monitor:write": true}) + ctx := authctx.WithPrincipal(context.Background(), principal) + authorize := mcpToolAuthorizer(db) + if err := authorize(ctx, builtin.ToolWaitToolExecution, map[string]interface{}{"execution_id": "exec-owned"}); err != nil { + t.Fatalf("owned execution denied: %v", err) + } + if err := authorize(ctx, builtin.ToolCancelToolExecution, map[string]interface{}{"execution_id": "exec-hidden"}); err == nil { + t.Fatal("foreign execution was allowed") + } +} + +func TestMCPAssetToolAuthorizationUsesAssetPermissionsAndScope(t *testing.T) { + db, err := database.NewDB(filepath.Join(t.TempDir(), "mcp-asset-authz.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + user, err := db.CreateRBACUser("asset-user", "Asset User", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + owned := &database.Asset{IP: "192.0.2.10", Port: 443, Protocol: "https"} + hidden := &database.Asset{IP: "192.0.2.20", Port: 443, Protocol: "https"} + if _, err := db.UpsertAssets([]*database.Asset{owned}, user.ID); err != nil { + t.Fatal(err) + } + if _, err := db.UpsertAssets([]*database.Asset{hidden}, ""); err != nil { + t.Fatal(err) + } + + permissions := map[string]bool{"asset:read": true, "asset:write": true} + ctx := authctx.WithPrincipal(context.Background(), authctx.NewPrincipal(user.ID, user.Username, database.RBACScopeAssigned, permissions)) + authorize := mcpToolAuthorizer(db) + if err := authorize(ctx, builtin.ToolQueryAssets, nil); err != nil { + t.Fatalf("asset query denied: %v", err) + } + if err := authorize(ctx, builtin.ToolGetAsset, map[string]interface{}{"id": owned.ID}); err != nil { + t.Fatalf("owned asset denied: %v", err) + } + if err := authorize(ctx, builtin.ToolGetAsset, map[string]interface{}{"id": hidden.ID}); err == nil { + t.Fatal("unassigned asset was readable") + } + if err := authorize(ctx, builtin.ToolUpdateAsset, map[string]interface{}{"id": owned.ID}); err != nil { + t.Fatalf("owned asset update denied: %v", err) + } + if err := authorize(ctx, builtin.ToolDeleteAsset, map[string]interface{}{"id": owned.ID}); err == nil { + t.Fatal("asset delete without asset:delete was allowed") + } +} + +func TestExternalMCPRequiresDedicatedPermission(t *testing.T) { + authorize := externalMCPToolAuthorizer() + ctx := authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("u1", "user", database.RBACScopeAssigned, map[string]bool{"agent:execute": true})) + if err := authorize(ctx, "server::tool", nil); err == nil { + t.Fatal("agent:execute alone authorized an external MCP tool") + } + ctx = authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("u1", "user", database.RBACScopeAll, map[string]bool{"mcp:external:execute": true})) + if err := authorize(ctx, "server::tool", nil); err != nil { + t.Fatalf("dedicated external MCP permission rejected: %v", err) + } +} + +func TestConfiguredCommandToolRequiresLocalExecutePermission(t *testing.T) { + authorize := mcpToolAuthorizer(nil) + agentOnly := authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("u1", "user", database.RBACScopeAssigned, map[string]bool{"agent:execute": true})) + if err := authorize(agentOnly, "nmap_scan", nil); err == nil { + t.Fatal("agent:execute alone authorized a configured command tool") + } + local := authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("u1", "user", database.RBACScopeAssigned, map[string]bool{"agent:local-execute": true})) + if err := authorize(local, "nmap_scan", nil); err != nil { + t.Fatalf("agent:local-execute rejected: %v", err) + } +} diff --git a/internal/app/mcp_http_auth_test.go b/internal/app/mcp_http_auth_test.go new file mode 100644 index 00000000..83ffbfc5 --- /dev/null +++ b/internal/app/mcp_http_auth_test.go @@ -0,0 +1,59 @@ +package app + +import ( + "bytes" + "net/http" + "net/http/httptest" + "path/filepath" + "testing" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/security" + + "go.uber.org/zap" +) + +func TestStandaloneMCPPrefersUserRBACAndDisablesGlobalTokenByDefault(t *testing.T) { + db, err := database.NewDB(filepath.Join(t.TempDir(), "mcp-http-auth.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + auth := security.NewAuthManager(12) + if _, err := auth.AttachRBACStore(db); err != nil { + t.Fatal(err) + } + hash, err := security.HashPassword("admin-secret") + if err != nil { + t.Fatal(err) + } + if err := db.UpdateRBACAdminPassword(hash); err != nil { + t.Fatal(err) + } + token, _, err := auth.Authenticate("admin", "admin-secret") + if err != nil { + t.Fatal(err) + } + server := mcp.NewServer(zap.NewNop()) + server.SetToolAuthorizer(mcpToolAuthorizer(db)) + a := &App{config: &config.Config{MCP: config.MCPConfig{AuthHeader: "X-MCP-Token", AuthHeaderValue: "static-secret"}}, auth: auth, mcpServer: server} + body := []byte(`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`) + + userReq := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewReader(body)) + userReq.Header.Set("Authorization", "Bearer "+token) + userW := httptest.NewRecorder() + a.mcpHandlerWithAuth(userW, userReq) + if userW.Code != http.StatusOK { + t.Fatalf("user bearer status = %d: %s", userW.Code, userW.Body.String()) + } + + staticReq := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewReader(body)) + staticReq.Header.Set("X-MCP-Token", "static-secret") + staticW := httptest.NewRecorder() + a.mcpHandlerWithAuth(staticW, staticReq) + if staticW.Code != http.StatusUnauthorized { + t.Fatalf("global static token status = %d, want 401", staticW.Code) + } +} diff --git a/internal/app/mcp_project_scope.go b/internal/app/mcp_project_scope.go new file mode 100644 index 00000000..171acbab --- /dev/null +++ b/internal/app/mcp_project_scope.go @@ -0,0 +1,26 @@ +package app + +import ( + "context" + "strings" + + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" +) + +func mcpEffectiveProjectFilter(ctx context.Context, db *database.DB) string { + if projectID := strings.TrimSpace(mcp.MCPProjectIDFromContext(ctx)); projectID != "" { + return projectID + } + if conversationID := mcpAuthorizationConversationID(ctx); conversationID != "" { + if db != nil { + if projectID, err := db.GetConversationProjectID(conversationID); err == nil { + if projectID = strings.TrimSpace(projectID); projectID != "" { + return projectID + } + } + } + return database.ProjectFilterUnbound + } + return "" +} diff --git a/internal/app/project_fact_tools.go b/internal/app/project_fact_tools.go new file mode 100644 index 00000000..2d570919 --- /dev/null +++ b/internal/app/project_fact_tools.go @@ -0,0 +1,389 @@ +package app + +import ( + "context" + "fmt" + "strings" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" + "cyberstrike-ai/internal/project" + + "go.uber.org/zap" +) + +func projectIDFromConversation(db *database.DB, ctx context.Context) (string, error) { + convID := agent.ConversationIDFromContext(ctx) + if convID == "" { + return "", fmt.Errorf("无法确定当前对话,请在对话上下文中使用项目事实工具") + } + pid, err := db.GetConversationProjectID(convID) + if err != nil { + return "", err + } + if strings.TrimSpace(pid) == "" { + return "", fmt.Errorf("当前对话未绑定项目,请先在对话中选择项目或创建带项目的对话") + } + return pid, nil +} + +func textResult(msg string, isErr bool) *mcp.ToolResult { + return &mcp.ToolResult{ + Content: []mcp.Content{{Type: "text", Text: msg}}, + IsError: isErr, + } +} + +// registerProjectFactTools 注册项目黑板 MCP 工具。 +func registerProjectFactTools(mcpServer *mcp.Server, db *database.DB, cfg *config.Config, logger *zap.Logger) { + if db == nil || cfg == nil || !cfg.Project.Enabled { + if logger != nil { + logger.Info("项目黑板工具未注册(未启用)") + } + return + } + + upsertTool := mcp.Tool{ + Name: builtin.ToolUpsertProjectFact, + Description: "写入或更新项目黑板事实,用于跨会话沉淀可复现上下文(非正式漏洞条目;可交付漏洞另用 record_vulnerability)。" + + "边渗透边记录:每确认新认知(端口/入口/凭据/可利用点)后立即调用,同 fact_key 覆盖更新,勿等会话结束。" + + "禁止仅写结论:summary 须含什么+在哪+如何验证;body 须含攻击链/请求响应/命令等复现细节。" + + "发现类建议 fact_key 为 finding|chain|exploit|poc/,category 对应 finding|chain|exploit|poc,body 按攻击链模板填写。" + + "环境类用 target|auth|infra|business/。同 fact_key 覆盖更新。需当前对话已绑定项目。", + ShortDescription: "写入/更新项目事实(含攻击链 body)", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "fact_key": map[string]interface{}{ + "type": "string", + "description": "项目内唯一 key:target/primary_domain、finding/sqli-login、exploit/upload-rce 等", + }, + "category": map[string]interface{}{ + "type": "string", + "description": "target | auth | infra | business | finding | chain | exploit | poc | note", + "enum": []string{"target", "auth", "infra", "business", "finding", "chain", "exploit", "poc", "note"}, + }, + "summary": map[string]interface{}{ + "type": "string", + "description": "索引用一行:结论 + 位置 + 触发/验证要点(勿仅写「存在 XSS」等空话)", + }, + "body": map[string]interface{}{ + "type": "string", + "description": "完整可复现详情(仅 get_project_fact 返回):须含攻击链步骤、原始 HTTP/命令、响应现象、证据与关联。" + + "发现/利用类首次写入必填;环境类建议含来源证据。攻击链类可参考模板章节:结论、目标与入口、攻击链、Exploit/POC、关键证据、关联、备注。" + + "更新已有 fact_key 时若省略或留空 body,将保留库中已有 body(可只改 summary)。", + }, + "confidence": map[string]interface{}{ + "type": "string", + "description": "confirmed | tentative | deprecated", + "enum": []string{"confirmed", "tentative", "deprecated"}, + }, + "pinned": map[string]interface{}{ + "type": "boolean", + "description": "是否优先出现在黑板索引", + }, + "related_vulnerability_id": map[string]interface{}{ + "type": "string", + "description": "可选:关联的漏洞记录 ID", + }, + "links": map[string]interface{}{ + "type": "array", + "description": "可选:关系边(from → 当前 fact)。finding 至少 1 条 {from:target/*, type:discovered_on};finding 上记录 exploit 用 {from:exploit/*, type:exploits}。省略保留已有边;传 [] 清空全部关系边。", + "items": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "from": map[string]interface{}{ + "type": "string", + "description": "来源 fact_key:存储为 from → 当前 fact", + }, + "type": map[string]interface{}{ + "type": "string", + "description": "depends_on | leads_to | enables | exploits | discovered_on | contains | part_of | supports", + }, + "confidence": map[string]interface{}{ + "type": "string", + "description": "confirmed | tentative | deprecated", + }, + }, + "required": []string{"from", "type"}, + }, + }, + }, + "required": []string{"fact_key", "summary"}, + }, + } + + mcpServer.RegisterTool(upsertTool, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + projectID, err := projectIDFromConversation(db, ctx) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + factKey, _ := args["fact_key"].(string) + summary, _ := args["summary"].(string) + if strings.TrimSpace(factKey) == "" || strings.TrimSpace(summary) == "" { + return textResult("错误: fact_key 与 summary 必填", true), nil + } + if len([]rune(summary)) > cfg.Project.FactSummaryMaxRunesEffective() { + return textResult(fmt.Sprintf("错误: summary 过长(最多 %d 字)", cfg.Project.FactSummaryMaxRunesEffective()), true), nil + } + f := &database.ProjectFact{ + ProjectID: projectID, + FactKey: factKey, + Category: strArg(args, "category"), + Summary: summary, + Body: strArg(args, "body"), + Confidence: strArg(args, "confidence"), + Pinned: boolArg(args, "pinned"), + RelatedVulnerabilityID: strArg(args, "related_vulnerability_id"), + } + if convID := agent.ConversationIDFromContext(ctx); convID != "" { + f.SourceConversationID = convID + } + created, err := db.UpsertProjectFact(f) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + if _, hasLinks := args["links"]; hasLinks { + linkInputs, err := project.ParseFactLinkInputs(args["links"]) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + convID := agent.ConversationIDFromContext(ctx) + if err := project.PersistFactLinksFromParsed(db, projectID, created.FactKey, convID, linkInputs, true); err != nil { + return textResult("错误: 保存关系边失败: "+err.Error(), true), nil + } + created, _ = db.GetProjectFactByKey(projectID, created.FactKey) + } else if parsed := project.ParseLinksFromBody(created.Body); len(parsed) > 0 { + if err := project.PersistFactIncomingLinks(db, projectID, created.FactKey, parsed, true); err != nil { + return textResult("错误: 从 body 解析边失败: "+err.Error(), true), nil + } + created, _ = db.GetProjectFactByKey(projectID, created.FactKey) + } + msg := fmt.Sprintf("事实已保存。\nfact_key: %s\nid: %s\nconfidence: %s", created.FactKey, created.ID, created.Confidence) + if in, _ := db.ListIncomingProjectFactEdges(projectID, created.FactKey); len(in) > 0 { + msg += "\n关系边: " + project.FormatFactLinksText(in) + } + if warn := project.SparseBodyWarningIfNeeded(f.Category, f.FactKey, f.Body); warn != "" { + msg += warn + } + return textResult(msg, false), nil + }) + + getTool := mcp.Tool{ + Name: builtin.ToolGetProjectFact, + Description: "按 fact_key 获取项目事实完整 body 与元数据。摘要不足时必须调用本工具,禁止臆造细节。", + ShortDescription: "按 key 获取事实详情", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "fact_key": map[string]interface{}{"type": "string", "description": "事实 key"}, + }, + "required": []string{"fact_key"}, + }, + } + mcpServer.RegisterTool(getTool, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + projectID, err := projectIDFromConversation(db, ctx) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + key := strings.TrimSpace(strArg(args, "fact_key")) + if key == "" { + return textResult("错误: fact_key 必填", true), nil + } + f, err := db.GetProjectFactByKey(projectID, key) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + msg := fmt.Sprintf("fact_key: %s\ncategory: %s\nconfidence: %s\nsummary: %s\nupdated_at: %s", + f.FactKey, f.Category, f.Confidence, f.Summary, f.UpdatedAt.Format("2006-01-02 15:04:05")) + if f.RelatedVulnerabilityID != "" { + msg += fmt.Sprintf("\nrelated_vulnerability_id: %s", f.RelatedVulnerabilityID) + } + if f.SourceConversationID != "" { + msg += fmt.Sprintf("\nsource_conversation_id: %s", f.SourceConversationID) + } + if in, _ := db.ListIncomingProjectFactEdges(projectID, f.FactKey); len(in) > 0 { + msg += "\n关系边(from → 本 fact):\n" + for _, e := range in { + msg += fmt.Sprintf("- %s ← %s (%s)\n", e.EdgeType, e.SourceFactKey, e.Confidence) + } + } + if out, _ := db.ListOutgoingProjectFactEdges(projectID, f.FactKey); len(out) > 0 { + msg += "指向其他事实:\n" + for _, e := range out { + msg += fmt.Sprintf("- %s → %s (%s)\n", e.EdgeType, e.TargetFactKey, e.Confidence) + } + } + msg += "\n\n--- body ---\n" + f.Body + if warn := project.SparseBodyWarningIfNeeded(f.Category, f.FactKey, f.Body); warn != "" { + msg += warn + } + return textResult(msg, false), nil + }) + + listTool := mcp.Tool{ + Name: builtin.ToolListProjectFacts, + Description: "列出当前项目的事实(分页)。", + ShortDescription: "列出项目事实", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "category": map[string]interface{}{"type": "string"}, + "confidence": map[string]interface{}{"type": "string"}, + "limit": map[string]interface{}{"type": "integer"}, + "offset": map[string]interface{}{"type": "integer"}, + }, + }, + } + mcpServer.RegisterTool(listTool, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + projectID, err := projectIDFromConversation(db, ctx) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + limit := intArg(args, "limit", 50) + offset := intArg(args, "offset", 0) + filter := database.ProjectFactListFilter{ + Category: strArg(args, "category"), + Confidence: strArg(args, "confidence"), + } + list, err := db.ListProjectFacts(projectID, filter, limit, offset) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + var b strings.Builder + b.WriteString(fmt.Sprintf("共 %d 条(limit=%d offset=%d):\n", len(list), limit, offset)) + for _, f := range list { + b.WriteString(fmt.Sprintf("- [%s] %s — %s (%s)\n", f.FactKey, f.Category, f.Summary, f.Confidence)) + } + return textResult(b.String(), false), nil + }) + + searchTool := mcp.Tool{ + Name: builtin.ToolSearchProjectFacts, + Description: "按关键词搜索项目事实(summary/body/fact_key)。", + ShortDescription: "搜索项目事实", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "query": map[string]interface{}{"type": "string"}, + "limit": map[string]interface{}{"type": "integer"}, + "offset": map[string]interface{}{"type": "integer"}, + }, + "required": []string{"query"}, + }, + } + mcpServer.RegisterTool(searchTool, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + projectID, err := projectIDFromConversation(db, ctx) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + q := strings.TrimSpace(strArg(args, "query")) + if q == "" { + return textResult("错误: query 必填", true), nil + } + list, err := db.ListProjectFacts(projectID, database.ProjectFactListFilter{Search: q}, intArg(args, "limit", 30), intArg(args, "offset", 0)) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + var b strings.Builder + b.WriteString(fmt.Sprintf("搜索 \"%s\" 命中 %d 条:\n", q, len(list))) + for _, f := range list { + b.WriteString(fmt.Sprintf("- [%s] %s — %s\n", f.FactKey, f.Category, f.Summary)) + } + return textResult(b.String(), false), nil + }) + + deprecateTool := mcp.Tool{ + Name: builtin.ToolDeprecateProjectFact, + Description: "将事实标记为 deprecated,从黑板索引中排除。", + ShortDescription: "废弃项目事实", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "fact_key": map[string]interface{}{"type": "string"}, + }, + "required": []string{"fact_key"}, + }, + } + mcpServer.RegisterTool(deprecateTool, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + projectID, err := projectIDFromConversation(db, ctx) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + key := strings.TrimSpace(strArg(args, "fact_key")) + if err := db.DeprecateProjectFact(projectID, key); err != nil { + return textResult("错误: "+err.Error(), true), nil + } + return textResult("事实已标记为 deprecated: "+key, false), nil + }) + + restoreTool := mcp.Tool{ + Name: builtin.ToolRestoreProjectFact, + Description: "将已废弃(deprecated)的事实恢复为 tentative 或 confirmed,重新参与黑板索引。", + ShortDescription: "恢复已废弃的项目事实", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "fact_key": map[string]interface{}{"type": "string"}, + "confidence": map[string]interface{}{ + "type": "string", + "description": "恢复后的置信度:tentative(默认)或 confirmed", + "enum": []string{"tentative", "confirmed"}, + }, + }, + "required": []string{"fact_key"}, + }, + } + mcpServer.RegisterTool(restoreTool, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + projectID, err := projectIDFromConversation(db, ctx) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + key := strings.TrimSpace(strArg(args, "fact_key")) + if key == "" { + return textResult("错误: fact_key 必填", true), nil + } + conf := strArg(args, "confidence") + if err := db.RestoreProjectFact(projectID, key, conf); err != nil { + return textResult("错误: "+err.Error(), true), nil + } + if conf == "" { + conf = "tentative" + } + return textResult(fmt.Sprintf("事实已恢复为 %s: %s", conf, key), false), nil + }) + + if logger != nil { + logger.Debug("项目黑板 MCP 工具注册成功") + } +} + +func strArg(args map[string]interface{}, key string) string { + if v, ok := args[key].(string); ok { + return v + } + return "" +} + +func boolArg(args map[string]interface{}, key string) bool { + if v, ok := args[key].(bool); ok { + return v + } + return false +} + +func intArg(args map[string]interface{}, key string, def int) int { + switch v := args[key].(type) { + case float64: + return int(v) + case int: + return v + case int64: + return int(v) + default: + return def + } +} diff --git a/internal/app/vision_tools.go b/internal/app/vision_tools.go new file mode 100644 index 00000000..f833588a --- /dev/null +++ b/internal/app/vision_tools.go @@ -0,0 +1,13 @@ +package app + +import ( + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/vision" + + "go.uber.org/zap" +) + +func registerVisionTools(mcpServer *mcp.Server, cfg *config.Config, logger *zap.Logger) { + vision.RegisterAnalyzeImageTool(mcpServer, cfg, logger) +} diff --git a/internal/app/vulnerability_tools.go b/internal/app/vulnerability_tools.go new file mode 100644 index 00000000..54d37819 --- /dev/null +++ b/internal/app/vulnerability_tools.go @@ -0,0 +1,466 @@ +package app + +import ( + "context" + "fmt" + "strings" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/authctx" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/mcp/builtin" + + "go.uber.org/zap" +) + +func conversationIDFromToolCtx(ctx context.Context) string { + if id := agent.ConversationIDFromContext(ctx); id != "" { + return id + } + return mcp.MCPConversationIDFromContext(ctx) +} + +// canAccessVulnerability 校验当前对话是否有权查看该漏洞(默认项目隔离,未绑项目则仅本会话)。 +func canAccessVulnerability(vuln *database.Vulnerability, convID, projectID string) bool { + if vuln == nil || convID == "" { + return false + } + if projectID != "" { + if strings.TrimSpace(vuln.ProjectID) == projectID { + return true + } + // 历史记录:写入时尚未绑定 project_id,但属于同一会话 + if strings.TrimSpace(vuln.ProjectID) == "" && vuln.ConversationID == convID { + return true + } + return false + } + return vuln.ConversationID == convID +} + +func buildVulnerabilityListFilter(db *database.DB, ctx context.Context, args map[string]interface{}) (database.VulnerabilityListFilter, string, error) { + convID := conversationIDFromToolCtx(ctx) + if convID == "" { + return database.VulnerabilityListFilter{}, "", fmt.Errorf("无法确定当前对话,请在对话上下文中使用漏洞查询工具") + } + + projectID := "" + if pid, err := db.GetConversationProjectID(convID); err == nil { + projectID = strings.TrimSpace(pid) + } + + scope := strings.TrimSpace(strArg(args, "scope")) + if scope == "" { + if projectID != "" { + scope = "project" + } else { + scope = "conversation" + } + } + + filter := database.VulnerabilityListFilter{ + Severity: strings.TrimSpace(strArg(args, "severity")), + Status: strings.TrimSpace(strArg(args, "status")), + } + if q := strings.TrimSpace(strArg(args, "q")); q != "" { + filter.Search = q + } else { + filter.Search = strings.TrimSpace(strArg(args, "search")) + } + + var scopeLabel string + switch scope { + case "project": + if projectID == "" { + return filter, "", fmt.Errorf("当前对话未绑定项目,无法按项目列出漏洞;请使用 scope=conversation,或先在对话中绑定项目") + } + filter.ProjectID = projectID + scopeLabel = fmt.Sprintf("项目 %s", projectID) + case "conversation": + filter.ConversationID = convID + scopeLabel = fmt.Sprintf("会话 %s", convID) + default: + return filter, "", fmt.Errorf("scope 仅支持 project 或 conversation,当前值: %s", scope) + } + return filter, scopeLabel, nil +} + +func formatVulnerabilityListItem(v *database.Vulnerability) string { + line := fmt.Sprintf("- id=%s | %s | %s | %s", v.ID, v.Severity, v.Status, v.Title) + if v.Type != "" { + line += fmt.Sprintf(" | type=%s", v.Type) + } + if v.Target != "" { + line += fmt.Sprintf(" | target=%s", truncateRunes(v.Target, 80)) + } + return line +} + +func formatVulnerabilityDetail(v *database.Vulnerability) string { + var b strings.Builder + b.WriteString(fmt.Sprintf("漏洞ID: %s\n", v.ID)) + b.WriteString(fmt.Sprintf("标题: %s\n", v.Title)) + b.WriteString(fmt.Sprintf("严重程度: %s\n", v.Severity)) + b.WriteString(fmt.Sprintf("状态: %s\n", v.Status)) + if v.Type != "" { + b.WriteString(fmt.Sprintf("类型: %s\n", v.Type)) + } + if v.Target != "" { + b.WriteString(fmt.Sprintf("目标: %s\n", v.Target)) + } + if v.ProjectID != "" { + b.WriteString(fmt.Sprintf("项目ID: %s\n", v.ProjectID)) + } + b.WriteString(fmt.Sprintf("会话ID: %s\n", v.ConversationID)) + if !v.CreatedAt.IsZero() { + b.WriteString(fmt.Sprintf("创建时间: %s\n", v.CreatedAt.Format("2006-01-02 15:04:05"))) + } + if v.Description != "" { + b.WriteString("\n--- 描述 ---\n") + b.WriteString(v.Description) + b.WriteString("\n") + } + if v.Preconditions != "" { + b.WriteString("\n--- 前置条件 ---\n") + b.WriteString(v.Preconditions) + b.WriteString("\n") + } + if v.ReproSteps != "" { + b.WriteString("\n--- 复现步骤 ---\n") + b.WriteString(v.ReproSteps) + b.WriteString("\n") + } + if v.Evidence != "" { + b.WriteString("\n--- 证据 / POC ---\n") + b.WriteString(v.Evidence) + b.WriteString("\n") + } + if v.Impact != "" { + b.WriteString("\n--- 影响 ---\n") + b.WriteString(v.Impact) + b.WriteString("\n") + } + if v.Recommendation != "" { + b.WriteString("\n--- 修复建议 ---\n") + b.WriteString(v.Recommendation) + b.WriteString("\n") + } + if v.RetestNotes != "" { + b.WriteString("\n--- 复测方式 ---\n") + b.WriteString(v.RetestNotes) + b.WriteString("\n") + } + return b.String() +} + +func missingVulnerabilityReproFields(args map[string]interface{}) []string { + required := []struct { + key string + label string + }{ + {"target", "target(受影响的 URL/IP/服务/接口)"}, + {"vulnerability_type", "vulnerability_type(漏洞类型)"}, + {"description", "description(漏洞摘要与触发点)"}, + {"reproduction_steps", "reproduction_steps(可逐步执行的复现步骤)"}, + {"evidence", "evidence(POC、原始请求/响应、命令输出或截图/日志证据)"}, + {"impact", "impact(确认后的实际影响)"}, + {"recommendation", "recommendation(修复建议)"}, + } + missing := make([]string, 0) + for _, item := range required { + if strings.TrimSpace(strArg(args, item.key)) == "" { + missing = append(missing, item.label) + } + } + return missing +} + +func truncateRunes(s string, max int) string { + r := []rune(s) + if len(r) <= max { + return s + } + return string(r[:max]) + "…" +} + +// registerVulnerabilityTools 注册漏洞记录与查询 MCP 工具。 +func registerVulnerabilityTools(mcpServer *mcp.Server, db *database.DB, logger *zap.Logger) { + registerRecordVulnerabilityTool(mcpServer, db, logger) + registerListVulnerabilitiesTool(mcpServer, db, logger) + registerGetVulnerabilityTool(mcpServer, db, logger) + if logger != nil { + logger.Debug("漏洞 MCP 工具注册成功", zap.Strings("tools", []string{ + builtin.ToolRecordVulnerability, + builtin.ToolListVulnerabilities, + builtin.ToolGetVulnerability, + })) + } +} + +func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, logger *zap.Logger) { + tool := mcp.Tool{ + Name: builtin.ToolRecordVulnerability, + Description: "记录发现的漏洞详情到漏洞管理系统。必须按“仅看本记录即可复现”的标准填写:目标、漏洞类型、触发点、复现步骤、证据/POC、实际影响和修复建议;前置条件与复测方式为推荐填写项。边渗透边记录:每验证出一条可复现漏洞后立即调用,勿等会话结束。记录前可先 list_vulnerabilities 避免重复。", + ShortDescription: "记录可复现的漏洞详情到漏洞管理系统", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "title": map[string]interface{}{ + "type": "string", + "description": "漏洞标题(必需)。建议格式:<资产/接口> 存在 <漏洞类型>,例如“/api/login 存在 SQL 注入”。", + }, + "description": map[string]interface{}{ + "type": "string", + "description": "漏洞摘要与触发点(必需):说明哪个功能/参数/入口存在问题、为什么可被利用。不要只写结论。", + }, + "severity": map[string]interface{}{ + "type": "string", + "description": "漏洞严重程度:critical(严重)、high(高)、medium(中)、low(低)、info(信息)", + "enum": []string{"critical", "high", "medium", "low", "info"}, + }, + "vulnerability_type": map[string]interface{}{ + "type": "string", + "description": "漏洞类型,如:SQL注入、XSS、CSRF、命令注入等(必需)", + }, + "target": map[string]interface{}{ + "type": "string", + "description": "受影响的目标(必需):尽量精确到 URL、IP:端口、服务名、接口路径和参数名。", + }, + "preconditions": map[string]interface{}{ + "type": "string", + "description": "前置条件(推荐填写):登录状态、权限、账号、Header/Cookie、特定数据、网络位置、环境/版本等;无前置条件可写“无”。", + }, + "reproduction_steps": map[string]interface{}{ + "type": "string", + "description": "复现步骤(必需):按 1/2/3 编号,写清入口、参数、payload、执行命令、观察点。应让未参与对话的人照做即可复现。", + }, + "evidence": map[string]interface{}{ + "type": "string", + "description": "证据 / POC(必需):原始 HTTP 请求/响应、curl/工具命令、截图文字说明、日志、DNSLog/回连记录、数据库结果、文件路径、时间戳等。优先放最小可验证证据。", + }, + "impact": map[string]interface{}{ + "type": "string", + "description": "漏洞影响说明(必需):结合已验证事实说明可造成什么后果,避免泛泛而谈。", + }, + "recommendation": map[string]interface{}{ + "type": "string", + "description": "修复建议(必需):给出针对该触发点/参数/组件的具体修复和复测建议。", + }, + "retest_notes": map[string]interface{}{ + "type": "string", + "description": "复测方式(推荐填写):修复后如何验证漏洞已关闭,包括应返回的状态码、错误信息或访问控制结果。", + }, + }, + "required": []string{"title", "description", "severity", "vulnerability_type", "target", "reproduction_steps", "evidence", "impact", "recommendation"}, + }, + } + + mcpServer.RegisterTool(tool, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + conversationID := strings.TrimSpace(strArg(args, "conversation_id")) + if conversationID == "" { + conversationID = conversationIDFromToolCtx(ctx) + } + if conversationID == "" { + return textResult("错误: conversation_id 未设置。这是系统错误,请重试。", true), nil + } + + title := strings.TrimSpace(strArg(args, "title")) + if title == "" { + return textResult("错误: title 参数必需且不能为空", true), nil + } + + severity := strings.TrimSpace(strArg(args, "severity")) + if severity == "" { + return textResult("错误: severity 参数必需且不能为空", true), nil + } + + validSeverities := map[string]bool{ + "critical": true, "high": true, "medium": true, "low": true, "info": true, + } + if !validSeverities[severity] { + return textResult(fmt.Sprintf("错误: severity 必须是 critical、high、medium、low 或 info 之一,当前值: %s", severity), true), nil + } + if missing := missingVulnerabilityReproFields(args); len(missing) > 0 { + return textResult("错误: 漏洞记录缺少必填信息,请补充后再记录:\n- "+strings.Join(missing, "\n- ")+"\n\n必填项用于确保单条记录可独立复现;前置条件和复测方式为推荐填写项。", true), nil + } + + projectID := "" + if pid, perr := db.GetConversationProjectID(conversationID); perr == nil { + projectID = strings.TrimSpace(pid) + } + + vuln := &database.Vulnerability{ + ConversationID: conversationID, + ProjectID: projectID, + Title: title, + Description: strArg(args, "description"), + Severity: severity, + Status: "open", + Type: strArg(args, "vulnerability_type"), + Target: strArg(args, "target"), + Preconditions: strArg(args, "preconditions"), + ReproSteps: strArg(args, "reproduction_steps"), + Evidence: strArg(args, "evidence"), + Impact: strArg(args, "impact"), + Recommendation: strArg(args, "recommendation"), + RetestNotes: strArg(args, "retest_notes"), + } + + created, err := db.CreateVulnerability(vuln) + if err != nil { + if logger != nil { + logger.Error("记录漏洞失败", zap.Error(err)) + } + return textResult(fmt.Sprintf("记录漏洞失败: %v", err), true), nil + } + if principal, ok := authctx.PrincipalFromContext(ctx); ok { + _ = db.SetResourceOwner("vulnerability", created.ID, principal.UserID) + _ = db.AssignResourceToUser(principal.UserID, "vulnerability", created.ID) + } + db.NotifyVulnerabilityCreated(created) + + if logger != nil { + logger.Info("漏洞记录成功", + zap.String("id", created.ID), + zap.String("title", created.Title), + zap.String("severity", created.Severity), + zap.String("conversation_id", conversationID), + ) + } + + return textResult(fmt.Sprintf("漏洞已成功记录!\n\n漏洞ID: %s\n标题: %s\n严重程度: %s\n状态: %s\n\n可使用 get_vulnerability(id) 查看详情,或 list_vulnerabilities 查看列表。", + created.ID, created.Title, created.Severity, created.Status), false), nil + }) +} + +func registerListVulnerabilitiesTool(mcpServer *mcp.Server, db *database.DB, logger *zap.Logger) { + tool := mcp.Tool{ + Name: builtin.ToolListVulnerabilities, + Description: "列出当前授权范围内的漏洞(摘要)。默认:对话已绑定项目时列出该项目下全部漏洞;未绑项目时仅列出当前会话漏洞。可用 scope=conversation 仅看本会话。记录新漏洞前建议先调用以避免重复。", + ShortDescription: "列出漏洞(默认当前项目)", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "scope": map[string]interface{}{ + "type": "string", + "description": "范围:project(默认,需绑定项目)| conversation(仅当前会话)", + "enum": []string{"project", "conversation"}, + }, + "severity": map[string]interface{}{ + "type": "string", + "description": "按严重程度筛选:critical、high、medium、low、info", + "enum": []string{"critical", "high", "medium", "low", "info"}, + }, + "status": map[string]interface{}{ + "type": "string", + "description": "按状态筛选:open、confirmed、fixed、false_positive、ignored", + "enum": []string{"open", "confirmed", "fixed", "false_positive", "ignored"}, + }, + "q": map[string]interface{}{ + "type": "string", + "description": "关键词搜索(标题、描述、类型、目标等)", + }, + "limit": map[string]interface{}{ + "type": "integer", + "description": "返回条数上限,默认 30,最大 100", + }, + "offset": map[string]interface{}{ + "type": "integer", + "description": "分页偏移,默认 0", + }, + }, + }, + } + + mcpServer.RegisterTool(tool, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + filter, scopeLabel, err := buildVulnerabilityListFilter(db, ctx, args) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + + limit := intArg(args, "limit", 30) + if limit <= 0 || limit > 100 { + limit = 30 + } + offset := intArg(args, "offset", 0) + if offset < 0 { + offset = 0 + } + + total, err := db.CountVulnerabilities(filter) + if err != nil { + if logger != nil { + logger.Warn("统计漏洞失败", zap.Error(err)) + } + total = 0 + } + + list, err := db.ListVulnerabilities(limit, offset, filter) + if err != nil { + return textResult("错误: "+err.Error(), true), nil + } + + var b strings.Builder + b.WriteString(fmt.Sprintf("范围: %s\n总计: %d | 本页: %d 条 (limit=%d offset=%d)\n\n", scopeLabel, total, len(list), limit, offset)) + if len(list) == 0 { + b.WriteString("(暂无漏洞记录)\n") + } else { + for _, v := range list { + b.WriteString(formatVulnerabilityListItem(v)) + b.WriteString("\n") + } + if total > offset+len(list) { + b.WriteString(fmt.Sprintf("\n(还有更多,可增大 offset 或使用 q/severity/status 筛选)\n")) + } + } + b.WriteString("\n需要 POC 与完整字段请对具体 id 调用 get_vulnerability。") + return textResult(b.String(), false), nil + }) +} + +func registerGetVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, logger *zap.Logger) { + tool := mcp.Tool{ + Name: builtin.ToolGetVulnerability, + Description: "按漏洞 ID 获取完整详情(含 POC、影响、修复建议)。仅能访问当前项目或当前会话下的漏洞(与 list_vulnerabilities 授权范围一致)。", + ShortDescription: "按 ID 获取漏洞详情", + InputSchema: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "id": map[string]interface{}{ + "type": "string", + "description": "漏洞 ID(list_vulnerabilities 返回的 id)", + }, + }, + "required": []string{"id"}, + }, + } + + mcpServer.RegisterTool(tool, func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) { + convID := conversationIDFromToolCtx(ctx) + if convID == "" { + return textResult("错误: 无法确定当前对话,请在对话上下文中使用本工具", true), nil + } + + id := strings.TrimSpace(strArg(args, "id")) + if id == "" { + return textResult("错误: id 必填", true), nil + } + + vuln, err := db.GetVulnerability(id) + if err != nil { + return textResult("错误: 漏洞不存在或查询失败", true), nil + } + + projectID := "" + if pid, perr := db.GetConversationProjectID(convID); perr == nil { + projectID = strings.TrimSpace(pid) + } + + if !canAccessVulnerability(vuln, convID, projectID) { + return textResult("错误: 无权访问该漏洞(仅可查看当前项目或当前会话下的记录)", true), nil + } + + return textResult(formatVulnerabilityDetail(vuln), false), nil + }) +} diff --git a/internal/einoobserve/attach.go b/internal/einoobserve/attach.go new file mode 100644 index 00000000..9846a2f7 --- /dev/null +++ b/internal/einoobserve/attach.go @@ -0,0 +1,455 @@ +// Package einoobserve attaches CloudWeGo Eino [callbacks.Handler] to ADK Runner contexts for +// structured logging and optional SSE trace events (eino_trace_*). +package einoobserve + +import ( + "context" + "encoding/json" + "fmt" + "strings" + "sync" + "sync/atomic" + "time" + + "cyberstrike-ai/internal/config" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/callbacks" + "github.com/cloudwego/eino/components" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" + "github.com/google/uuid" + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/codes" + "go.opentelemetry.io/otel/trace" + "go.uber.org/zap" +) + +type ctxSpanKey struct{} + +type ctxOtelSpanKey struct{} + +// Params for attaching per-run callback instrumentation. +type Params struct { + Logger *zap.Logger + Progress func(eventType, message string, data interface{}) + ConversationID string + OrchMode string + OrchestratorName string + RunID string +} + +// AttachAgentRunCallbacks returns ctx wrapped with callbacks.InitCallbacks when enabled. +// Safe to call with nil cfg or disabled cfg (returns ctx unchanged). +func AttachAgentRunCallbacks(ctx context.Context, cfg *config.MultiAgentEinoCallbacksConfig, p Params) context.Context { + if ctx == nil { + return ctx + } + if cfg == nil || !cfg.Enabled { + return ctx + } + mode := cfg.EinoCallbacksModeEffective() + if mode == "off" { + return ctx + } + runID := strings.TrimSpace(p.RunID) + if runID == "" { + runID = uuid.New().String() + } + if p.Progress != nil && cfg.ShouldEmitEinoTraceSSE(mode) { + p.Progress("eino_trace_run", "Eino callbacks session", map[string]interface{}{ + "runId": runID, + "conversationId": strings.TrimSpace(p.ConversationID), + "orchestration": strings.TrimSpace(p.OrchMode), + "orchestratorName": strings.TrimSpace(p.OrchestratorName), + "observeMode": mode, + "source": "eino_callbacks", + }) + } + h := &runHandler{ + cfg: *cfg, + mode: mode, + params: p, + runID: runID, + } + b := callbacks.NewHandlerBuilder(). + OnStartFn(h.onStart). + OnEndFn(h.onEnd). + OnErrorFn(h.onError) + if mode == "full" { + b = b.OnStartWithStreamInputFn(h.onStartStreamIn).OnEndWithStreamOutputFn(h.onEndStreamOut) + } + ri := &callbacks.RunInfo{ + Name: "CyberStrikeADKRun", + Type: strings.TrimSpace(p.OrchMode), + Component: components.Component("AgentSession"), + } + return callbacks.InitCallbacks(ctx, ri, b.Build()) +} + +type runHandler struct { + cfg config.MultiAgentEinoCallbacksConfig + mode string + params Params + runID string + + mu sync.Mutex + spanStack []string + seq atomic.Uint64 +} + +func safeRunInfo(info *callbacks.RunInfo) callbacks.RunInfo { + if info == nil { + return callbacks.RunInfo{ + Name: "unknown", + Type: "unknown", + Component: components.Component("unknown"), + } + } + return *info +} + +func (h *runHandler) genSpanID() string { + return fmt.Sprintf("%s-%d", h.runID, h.seq.Add(1)) +} + +func (h *runHandler) popSpan() (id string) { + h.mu.Lock() + defer h.mu.Unlock() + if len(h.spanStack) == 0 { + return "" + } + id = h.spanStack[len(h.spanStack)-1] + h.spanStack = h.spanStack[:len(h.spanStack)-1] + return id +} + +// popMatching removes the given id from the stack top if it matches; otherwise pops until empty or match (rare ordering mismatch). +func (h *runHandler) popMatching(want string) string { + h.mu.Lock() + defer h.mu.Unlock() + if want == "" { + if len(h.spanStack) == 0 { + return "" + } + id := h.spanStack[len(h.spanStack)-1] + h.spanStack = h.spanStack[:len(h.spanStack)-1] + return id + } + for len(h.spanStack) > 0 { + top := h.spanStack[len(h.spanStack)-1] + h.spanStack = h.spanStack[:len(h.spanStack)-1] + if top == want { + return top + } + } + return want +} + +func (h *runHandler) onStart(ctx context.Context, info *callbacks.RunInfo, input callbacks.CallbackInput) context.Context { + ri := safeRunInfo(info) + var parentID string + h.mu.Lock() + if len(h.spanStack) > 0 { + parentID = h.spanStack[len(h.spanStack)-1] + } + spanID := h.genSpanID() + h.spanStack = append(h.spanStack, spanID) + h.mu.Unlock() + + inSum := summarizeCallbackInput(input, h.cfg.EinoCallbacksMaxInputSummaryRunes()) + if h.cfg.OtelTracingActive() { + tracer := otel.Tracer("cyberstrike/eino") + spanName := callbackSpanName(info) + var sp trace.Span + ctx, sp = tracer.Start(ctx, spanName, + trace.WithSpanKind(trace.SpanKindInternal), + trace.WithAttributes( + attribute.String("eino.component", string(ri.Component)), + attribute.String("eino.name", ri.Name), + attribute.String("eino.type", ri.Type), + attribute.String("cyberstrike.run_id", h.runID), + attribute.String("cyberstrike.conversation_id", strings.TrimSpace(h.params.ConversationID)), + attribute.String("cyberstrike.orchestration", strings.TrimSpace(h.params.OrchMode)), + ), + ) + if inSum != "" { + sp.SetAttributes(attribute.String("eino.input.summary", truncateForAttr(inSum, 256))) + } + ctx = context.WithValue(ctx, ctxOtelSpanKey{}, sp) + } + if h.params.Logger != nil { + fields := []zap.Field{ + zap.String("runId", h.runID), + zap.String("spanId", spanID), + zap.String("parentSpanId", parentID), + zap.String("component", string(ri.Component)), + zap.String("name", ri.Name), + zap.String("type", ri.Type), + zap.String("phase", "start"), + } + if sp, ok := ctx.Value(ctxOtelSpanKey{}).(trace.Span); ok && sp != nil { + if sc := sp.SpanContext(); sc.IsValid() { + fields = append(fields, + zap.String("trace_id", sc.TraceID().String()), + zap.String("otel_span_id", sc.SpanID().String()), + ) + } + } + if h.cfg.ZapVerbose { + h.params.Logger.Debug("eino_callback", append(fields, zap.String("inputSummary", inSum))...) + } else { + h.params.Logger.Info("eino_callback", fields...) + } + } + if h.params.Progress != nil && h.cfg.ShouldEmitEinoTraceSSE(h.mode) { + h.params.Progress("eino_trace_start", "", map[string]interface{}{ + "runId": h.runID, + "spanId": spanID, + "parentSpanId": parentID, + "conversationId": strings.TrimSpace(h.params.ConversationID), + "orchestration": strings.TrimSpace(h.params.OrchMode), + "component": string(ri.Component), + "name": ri.Name, + "type": ri.Type, + "ts": time.Now().UTC().Format(time.RFC3339Nano), + "inputSummary": inSum, + "source": "eino_callbacks", + }) + } + ctx = context.WithValue(ctx, ctxSpanKey{}, spanID) + return ctx +} + +func (h *runHandler) onEnd(ctx context.Context, info *callbacks.RunInfo, output callbacks.CallbackOutput) context.Context { + ri := safeRunInfo(info) + spanID, _ := ctx.Value(ctxSpanKey{}).(string) + if spanID == "" { + spanID = h.popSpan() + } else { + spanID = h.popMatching(spanID) + } + outSum := summarizeCallbackOutput(output, h.cfg.EinoCallbacksMaxOutputSummaryRunes()) + if sp, ok := ctx.Value(ctxOtelSpanKey{}).(trace.Span); ok && sp != nil { + if outSum != "" { + sp.SetAttributes(attribute.String("eino.output.summary", truncateForAttr(outSum, 256))) + } + sp.SetStatus(codes.Ok, "") + sp.End() + } + if h.params.Logger != nil { + fields := []zap.Field{ + zap.String("runId", h.runID), + zap.String("spanId", spanID), + zap.String("component", string(ri.Component)), + zap.String("name", ri.Name), + zap.String("type", ri.Type), + zap.String("phase", "end"), + } + if h.cfg.ZapVerbose { + h.params.Logger.Debug("eino_callback", append(fields, zap.String("outputSummary", outSum))...) + } else { + h.params.Logger.Info("eino_callback", fields...) + } + } + if h.params.Progress != nil && h.cfg.ShouldEmitEinoTraceSSE(h.mode) { + h.params.Progress("eino_trace_end", "", map[string]interface{}{ + "runId": h.runID, + "spanId": spanID, + "conversationId": strings.TrimSpace(h.params.ConversationID), + "orchestration": strings.TrimSpace(h.params.OrchMode), + "component": string(ri.Component), + "name": ri.Name, + "type": ri.Type, + "ts": time.Now().UTC().Format(time.RFC3339Nano), + "outputSummary": outSum, + "source": "eino_callbacks", + }) + } + return ctx +} + +func (h *runHandler) onError(ctx context.Context, info *callbacks.RunInfo, err error) context.Context { + ri := safeRunInfo(info) + spanID, _ := ctx.Value(ctxSpanKey{}).(string) + if spanID == "" { + spanID = h.popSpan() + } else { + spanID = h.popMatching(spanID) + } + msg := "" + if err != nil { + msg = truncateRunes(err.Error(), h.cfg.EinoCallbacksMaxOutputSummaryRunes()) + } + if sp, ok := ctx.Value(ctxOtelSpanKey{}).(trace.Span); ok && sp != nil { + if err != nil { + sp.RecordError(err) + } + sp.SetStatus(codes.Error, msg) + sp.End() + } + if h.params.Logger != nil { + h.params.Logger.Warn("eino_callback_error", + zap.String("runId", h.runID), + zap.String("spanId", spanID), + zap.String("component", string(ri.Component)), + zap.String("name", ri.Name), + zap.String("type", ri.Type), + zap.Error(err), + ) + } + if h.params.Progress != nil && h.cfg.ShouldEmitEinoTraceSSE(h.mode) { + h.params.Progress("eino_trace_error", msg, map[string]interface{}{ + "runId": h.runID, + "spanId": spanID, + "conversationId": strings.TrimSpace(h.params.ConversationID), + "orchestration": strings.TrimSpace(h.params.OrchMode), + "component": string(ri.Component), + "name": ri.Name, + "type": ri.Type, + "ts": time.Now().UTC().Format(time.RFC3339Nano), + "error": msg, + "source": "eino_callbacks", + }) + } + return ctx +} + +func (h *runHandler) onStartStreamIn(ctx context.Context, info *callbacks.RunInfo, input *schema.StreamReader[callbacks.CallbackInput]) context.Context { + ri := safeRunInfo(info) + if input != nil { + input.Close() + } + if h.params.Logger != nil { + h.params.Logger.Debug("eino_callback_stream_in", + zap.String("runId", h.runID), + zap.String("component", string(ri.Component)), + zap.String("name", ri.Name), + ) + } + return ctx +} + +func (h *runHandler) onEndStreamOut(ctx context.Context, info *callbacks.RunInfo, output *schema.StreamReader[callbacks.CallbackOutput]) context.Context { + ri := safeRunInfo(info) + if output != nil { + output.Close() + } + if h.params.Logger != nil { + h.params.Logger.Debug("eino_callback_stream_out", + zap.String("runId", h.runID), + zap.String("component", string(ri.Component)), + zap.String("name", ri.Name), + ) + } + return ctx +} + +func callbackSpanName(info *callbacks.RunInfo) string { + if info == nil { + return "eino.callback" + } + comp := strings.TrimSpace(string(info.Component)) + name := strings.TrimSpace(info.Name) + typ := strings.TrimSpace(info.Type) + if name != "" && comp != "" { + return comp + "/" + name + } + if typ != "" && comp != "" { + return comp + "[" + typ + "]" + } + if comp != "" { + return comp + } + return "eino.callback" +} + +func truncateForAttr(s string, maxRunes int) string { + return truncateRunes(s, maxRunes) +} + +func summarizeCallbackInput(in callbacks.CallbackInput, maxRunes int) string { + if in == nil { + return "" + } + if ai := adk.ConvAgentCallbackInput(in); ai != nil { + parts := []string{"agent"} + if ai.Input != nil { + parts = append(parts, fmt.Sprintf("messages=%d", len(ai.Input.Messages))) + } + if ai.ResumeInfo != nil { + parts = append(parts, "resume=true") + } + return strings.Join(parts, " ") + } + if mi := model.ConvCallbackInput(in); mi != nil { + return fmt.Sprintf("chatModel messages=%d tools=%d", len(mi.Messages), len(mi.Tools)) + } + if ti := tool.ConvCallbackInput(in); ti != nil { + raw := ti.ArgumentsInJSON + return "tool args=" + truncateRunes(raw, maxRunes) + } + b, err := json.Marshal(in) + if err != nil { + return fmt.Sprintf("%T", in) + } + return truncateRunes(string(b), maxRunes) +} + +func summarizeCallbackOutput(out callbacks.CallbackOutput, maxRunes int) string { + if out == nil { + return "" + } + if ao := adk.ConvAgentCallbackOutput(out); ao != nil { + return "agent_events=stream" + } + if mo := model.ConvCallbackOutput(out); mo != nil && mo.Message != nil { + s := "" + if mo.Message.Content != "" { + s = mo.Message.Content + } + if mo.TokenUsage != nil { + return fmt.Sprintf("tokens total=%d completion=%d prompt=%d text=%s", + mo.TokenUsage.TotalTokens, mo.TokenUsage.CompletionTokens, mo.TokenUsage.PromptTokens, + truncateRunes(s, minInt(120, maxRunes))) + } + return "assistant len=" + itoa(len(s)) + } + if to := tool.ConvCallbackOutput(out); to != nil { + if to.Response != "" { + return truncateRunes(to.Response, maxRunes) + } + if to.ToolOutput != nil { + return "tool_result multimodal" + } + } + b, err := json.Marshal(out) + if err != nil { + return fmt.Sprintf("%T", out) + } + return truncateRunes(string(b), maxRunes) +} + +func minInt(a, b int) int { + if a < b { + return a + } + return b +} + +func itoa(n int) string { + return fmt.Sprintf("%d", n) +} + +func truncateRunes(s string, maxRunes int) string { + if maxRunes <= 0 { + return "" + } + r := []rune(s) + if len(r) <= maxRunes { + return s + } + return string(r[:maxRunes]) + "…" +} diff --git a/internal/einoobserve/attach_test.go b/internal/einoobserve/attach_test.go new file mode 100644 index 00000000..d12290a2 --- /dev/null +++ b/internal/einoobserve/attach_test.go @@ -0,0 +1,49 @@ +package einoobserve + +import ( + "context" + "testing" + + "cyberstrike-ai/internal/config" +) + +func TestAttachAgentRunCallbacks_Disabled(t *testing.T) { + ctx := context.Background() + cfg := &config.MultiAgentEinoCallbacksConfig{Enabled: false} + out := AttachAgentRunCallbacks(ctx, cfg, Params{}) + if out != ctx { + t.Fatalf("expected same ctx when disabled") + } +} + +func TestAttachAgentRunCallbacksUsesProvidedRunID(t *testing.T) { + emit := true + var gotRunID string + ctx := context.Background() + cfg := &config.MultiAgentEinoCallbacksConfig{Enabled: true, Mode: "sse", SseTraceToClient: &emit} + + AttachAgentRunCallbacks(ctx, cfg, Params{ + RunID: "run-shared", + Progress: func(eventType, _ string, data interface{}) { + if eventType != "eino_trace_run" { + return + } + if m, ok := data.(map[string]interface{}); ok { + gotRunID, _ = m["runId"].(string) + } + }, + }) + + if gotRunID != "run-shared" { + t.Fatalf("runId = %q, want run-shared", gotRunID) + } +} + +func TestTruncateRunes(t *testing.T) { + if got := truncateRunes("abc", 10); got != "abc" { + t.Fatalf("got %q", got) + } + if got := truncateRunes("abcdefghij", 4); got != "abcd…" { + t.Fatalf("got %q", got) + } +} diff --git a/internal/einoobserve/otel.go b/internal/einoobserve/otel.go new file mode 100644 index 00000000..05800abd --- /dev/null +++ b/internal/einoobserve/otel.go @@ -0,0 +1,111 @@ +package einoobserve + +import ( + "context" + "fmt" + "strings" + "sync" + + "cyberstrike-ai/internal/config" + + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp" + "go.opentelemetry.io/otel/exporters/stdout/stdouttrace" + "go.opentelemetry.io/otel/sdk/resource" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + semconv "go.opentelemetry.io/otel/semconv/v1.26.0" + "go.uber.org/zap" +) + +var ( + otelMu sync.Mutex + otelShutdown func(context.Context) error + otelInitialized bool +) + +// InitOtelFromConfig installs the global OpenTelemetry TracerProvider when +// eino_callbacks.otel is enabled and exporter is not none. Safe to call multiple times. +func InitOtelFromConfig(cfg *config.MultiAgentEinoCallbacksConfig, log *zap.Logger) (shutdown func(context.Context) error, err error) { + shutdown = func(context.Context) error { return nil } + if cfg == nil || !cfg.OtelTracingActive() { + return shutdown, nil + } + + otelMu.Lock() + defer otelMu.Unlock() + if otelInitialized { + if otelShutdown != nil { + return otelShutdown, nil + } + return shutdown, nil + } + + oc := cfg.Otel + expKind := oc.OtelExporterEffective() + ctx := context.Background() + + var exporter sdktrace.SpanExporter + switch expKind { + case "stdout": + exporter, err = stdouttrace.New() + if err != nil { + return shutdown, fmt.Errorf("eino otel stdout exporter: %w", err) + } + case "otlphttp": + ep := strings.TrimSpace(oc.OTLPEndpoint) + if ep == "" { + ep = "localhost:4318" + } + exporter, err = otlptracehttp.New(ctx, + otlptracehttp.WithEndpoint(ep), + otlptracehttp.WithURLPath("/v1/traces"), + ) + if err != nil { + return shutdown, fmt.Errorf("eino otel otlphttp exporter: %w", err) + } + default: + return shutdown, nil + } + + res, err := resource.New(ctx, + resource.WithAttributes( + semconv.ServiceName(oc.ServiceNameEffective()), + ), + ) + if err != nil { + return shutdown, fmt.Errorf("eino otel resource: %w", err) + } + + sampler := sdktrace.ParentBased(sdktrace.TraceIDRatioBased(oc.SampleRatioEffective())) + tp := sdktrace.NewTracerProvider( + sdktrace.WithBatcher(exporter), + sdktrace.WithResource(res), + sdktrace.WithSampler(sampler), + ) + otel.SetTracerProvider(tp) + + otelShutdown = tp.Shutdown + otelInitialized = true + if log != nil { + log.Info("eino otel: tracer provider initialized", + zap.String("exporter", expKind), + zap.String("service", oc.ServiceNameEffective()), + zap.Float64("sample_ratio", oc.SampleRatioEffective()), + ) + } + return otelShutdown, nil +} + +// ShutdownOtel flushes and shuts down the global TracerProvider if it was installed. +func ShutdownOtel(ctx context.Context) error { + otelMu.Lock() + fn := otelShutdown + otelShutdown = nil + inited := otelInitialized + otelInitialized = false + otelMu.Unlock() + if !inited || fn == nil { + return nil + } + return fn(ctx) +} diff --git a/internal/openai/claude_bridge.go b/internal/openai/claude_bridge.go new file mode 100644 index 00000000..530e9f9b --- /dev/null +++ b/internal/openai/claude_bridge.go @@ -0,0 +1,1293 @@ +package openai + +// claude_bridge.go 将 OpenAI 格式的请求/响应自动转换为 Anthropic Claude Messages API 格式。 +// 当 config.Provider == "claude" 时,Client 自动走此桥接层,对上层调用方完全透明。 +// +// 转换规则: +// Request: OpenAI /chat/completions → Claude /v1/messages +// Response: Claude /v1/messages → OpenAI /chat/completions 格式 +// Stream: Claude SSE (event: content_block_delta / message_delta) → OpenAI SSE 格式 +// Auth: Bearer → x-api-key +// Tools: OpenAI tools[] → Claude tools[] (input_schema) +// +// Extended thinking: 顶层 `thinking` / `output_config` 从 OpenAI 请求体透传;响应中 `thinking` block 映射为 +// `reasoning_content`(可读前缀 + 内部 JSON 尾缀以保留 signature,供多轮工具续跑;UI 用 openai.DisplayReasoningContent 剥离)。 + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "cyberstrike-ai/internal/config" + + "go.uber.org/zap" +) + +// ============================================================ +// Claude Request Types +// ============================================================ + +// claudeRequest 表示 Anthropic Messages API 的请求体。 +type claudeRequest struct { + Model string `json:"model"` + MaxTokens int `json:"max_tokens"` + System string `json:"system,omitempty"` + Messages []claudeMessage `json:"messages"` + Tools []claudeTool `json:"tools,omitempty"` + Stream bool `json:"stream,omitempty"` + Thinking json.RawMessage `json:"thinking,omitempty"` + OutputConfig json.RawMessage `json:"output_config,omitempty"` +} + +type claudeMessage struct { + Role string `json:"role"` + Content claudeMessageContent `json:"content"` +} + +// claudeMessageContent 可以是纯字符串或 content block 数组。 +// MarshalJSON / UnmarshalJSON 自动处理两种形式。 +type claudeMessageContent struct { + Text string // 纯文本形式(简写) + Blocks []claudeContentBlock // 多 block 形式(tool_use / tool_result 必须用这种) +} + +func (c claudeMessageContent) MarshalJSON() ([]byte, error) { + if len(c.Blocks) > 0 { + return json.Marshal(c.Blocks) + } + return json.Marshal(c.Text) +} + +func (c *claudeMessageContent) UnmarshalJSON(data []byte) error { + // 尝试字符串 + var s string + if err := json.Unmarshal(data, &s); err == nil { + c.Text = s + return nil + } + // 尝试数组 + return json.Unmarshal(data, &c.Blocks) +} + +type claudeContentBlock struct { + Type string `json:"type"` + + // text block + Text string `json:"text,omitempty"` + + // thinking block (extended thinking) + Thinking string `json:"thinking,omitempty"` + Signature string `json:"signature,omitempty"` + + // tool_use block (assistant 返回) + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + Input json.RawMessage `json:"input,omitempty"` + + // tool_result block (user 提交) + ToolUseID string `json:"tool_use_id,omitempty"` + Content string `json:"content,omitempty"` + IsError bool `json:"is_error,omitempty"` +} + +type claudeTool struct { + Name string `json:"name"` + Description string `json:"description,omitempty"` + InputSchema map[string]interface{} `json:"input_schema"` +} + +// ============================================================ +// Claude Response Types +// ============================================================ + +type claudeResponse struct { + ID string `json:"id"` + Type string `json:"type"` + Role string `json:"role"` + Content []claudeContentBlock `json:"content"` + Model string `json:"model"` + StopReason string `json:"stop_reason"` + StopSequence *string `json:"stop_sequence"` + Usage *claudeUsage `json:"usage,omitempty"` + Error *claudeError `json:"error,omitempty"` +} + +type claudeUsage struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` +} + +type claudeError struct { + Type string `json:"type"` + Message string `json:"message"` +} + +// ============================================================ +// Conversion: OpenAI Request → Claude Request +// ============================================================ + +// convertOpenAIToClaude 将任意 OpenAI payload (map 或 struct) 转换为 claudeRequest。 +func convertOpenAIToClaude(payload interface{}) (*claudeRequest, error) { + // 先统一序列化为 JSON,再以 map 反序列化,方便处理各种输入形式 + raw, err := json.Marshal(payload) + if err != nil { + return nil, fmt.Errorf("claude bridge: marshal payload: %w", err) + } + + var oai map[string]interface{} + if err := json.Unmarshal(raw, &oai); err != nil { + return nil, fmt.Errorf("claude bridge: unmarshal payload: %w", err) + } + + req := &claudeRequest{} + + // model + if m, ok := oai["model"].(string); ok { + req.Model = m + } + + // Anthropic requires max_tokens. OpenAI-compatible clients prefer + // max_completion_tokens, so map it first and keep max_tokens as fallback. + if mt, ok := oai["max_completion_tokens"].(float64); ok && mt > 0 { + req.MaxTokens = int(mt) + } else if mt, ok := oai["max_tokens"].(float64); ok && mt > 0 { + req.MaxTokens = int(mt) + } else { + req.MaxTokens = 8192 // Claude 默认最大输出(兼容 Haiku/Sonnet/Opus) + } + + // stream + if s, ok := oai["stream"].(bool); ok { + req.Stream = s + } + + // messages + msgs, _ := oai["messages"].([]interface{}) + for i := 0; i < len(msgs); i++ { + mm, ok := msgs[i].(map[string]interface{}) + if !ok { + continue + } + role, _ := mm["role"].(string) + content, _ := mm["content"].(string) + + // system message → 提取到顶级 system 字段 + if role == "system" { + if req.System != "" { + req.System += "\n\n" + } + req.System += content + continue + } + + // tool_calls (assistant 消息中包含工具调用) + if role == "assistant" { + rc, _ := mm["reasoning_content"].(string) + _, thinkingReplay := parseClaudeReasoningAssistantBlocks(rc) + + var blocks []claudeContentBlock + for _, tb := range thinkingReplay { + blocks = append(blocks, tb) + } + if content != "" { + blocks = append(blocks, claudeContentBlock{Type: "text", Text: content}) + } + + if tcs, ok := mm["tool_calls"].([]interface{}); ok { + for _, tc := range tcs { + tcMap, ok := tc.(map[string]interface{}) + if !ok { + continue + } + tcID, _ := tcMap["id"].(string) + fn, _ := tcMap["function"].(map[string]interface{}) + fnName, _ := fn["name"].(string) + fnArgs, _ := fn["arguments"] + + // 防御:缺少 name 或 id 的 tool_call 会被 Claude 拒绝 + if strings.TrimSpace(fnName) == "" { + fnName = "unknown_function" + } + if strings.TrimSpace(tcID) == "" { + tcID = fmt.Sprintf("call_%d", time.Now().UnixNano()) + } + + var inputRaw json.RawMessage + switch v := fnArgs.(type) { + case string: + inputRaw = json.RawMessage(v) + default: + inputRaw, _ = json.Marshal(v) + } + // 防止空字符串/非法 JSON 导致 Marshal 失败 + if len(inputRaw) == 0 || !json.Valid(inputRaw) { + inputRaw = json.RawMessage("{}") + } + blocks = append(blocks, claudeContentBlock{ + Type: "tool_use", + ID: tcID, + Name: fnName, + Input: inputRaw, + }) + } + } + + if len(blocks) > 0 { + req.Messages = append(req.Messages, claudeMessage{ + Role: "assistant", + Content: claudeMessageContent{Blocks: blocks}, + }) + } + continue + } + + // tool result (role == "tool" in OpenAI) + // Claude 要求同一轮的多个 tool_result 合并为一个 user 消息(多 block), + // 否则违反 user/assistant 交替规则。 + if role == "tool" { + var toolBlocks []claudeContentBlock + // 收集当前及后续连续的 tool 消息 + for ; i < len(msgs); i++ { + tmm, ok := msgs[i].(map[string]interface{}) + if !ok { + break + } + tr, _ := tmm["role"].(string) + if tr != "tool" { + break + } + tcID, _ := tmm["tool_call_id"].(string) + tcContent, _ := tmm["content"].(string) + toolBlocks = append(toolBlocks, claudeContentBlock{ + Type: "tool_result", + ToolUseID: tcID, + Content: tcContent, + }) + } + i-- // 外层 for 会 i++,回退一步 + req.Messages = append(req.Messages, claudeMessage{ + Role: "user", + Content: claudeMessageContent{Blocks: toolBlocks}, + }) + continue + } + + // 普通 user/assistant 消息 + req.Messages = append(req.Messages, claudeMessage{ + Role: role, + Content: claudeMessageContent{Text: content}, + }) + } + + // tools + if tools, ok := oai["tools"].([]interface{}); ok { + for _, t := range tools { + tMap, ok := t.(map[string]interface{}) + if !ok { + continue + } + fn, ok := tMap["function"].(map[string]interface{}) + if !ok { + continue + } + ct := claudeTool{} + ct.Name, _ = fn["name"].(string) + ct.Description, _ = fn["description"].(string) + if params, ok := fn["parameters"].(map[string]interface{}); ok { + ct.InputSchema = params + } else { + ct.InputSchema = map[string]interface{}{"type": "object", "properties": map[string]interface{}{}} + } + req.Tools = append(req.Tools, ct) + } + } + + // Extended thinking + effort (Anthropic top-level); merged from Eino ExtraFields / admin extras. + if th, ok := oai["thinking"]; ok && th != nil { + if raw, err := json.Marshal(th); err == nil && len(raw) > 0 && string(raw) != "null" { + req.Thinking = json.RawMessage(raw) + } + } + if oc, ok := oai["output_config"]; ok && oc != nil { + if raw, err := json.Marshal(oc); err == nil && len(raw) > 0 && string(raw) != "null" { + req.OutputConfig = json.RawMessage(raw) + } + } + if err := validateClaudeToolPairs(req.Messages); err != nil { + return nil, err + } + + return req, nil +} + +// validateClaudeToolPairs prevents malformed OpenAI history from reaching +// Anthropic/Bedrock. Every tool_use batch must be answered by the immediately +// following user message, with exactly one tool_result for every ID. +func validateClaudeToolPairs(messages []claudeMessage) error { + validatedResultMessages := make(map[int]struct{}) + for i, msg := range messages { + expected := make(map[string]struct{}) + for _, block := range msg.Content.Blocks { + if block.Type != "tool_use" { + continue + } + id := strings.TrimSpace(block.ID) + if id == "" { + return fmt.Errorf("claude bridge: assistant message %d has tool_use without id", i) + } + if _, duplicate := expected[id]; duplicate { + return fmt.Errorf("claude bridge: assistant message %d has duplicate tool_use id %q", i, id) + } + expected[id] = struct{}{} + } + if len(expected) == 0 { + continue + } + if i+1 >= len(messages) || messages[i+1].Role != "user" { + return fmt.Errorf("claude bridge: assistant message %d tool_use is not immediately followed by user tool_result", i) + } + seen := make(map[string]struct{}, len(expected)) + for _, block := range messages[i+1].Content.Blocks { + if block.Type != "tool_result" { + continue + } + id := strings.TrimSpace(block.ToolUseID) + if _, ok := expected[id]; !ok { + return fmt.Errorf("claude bridge: user message %d has unexpected tool_result id %q", i+1, id) + } + if _, duplicate := seen[id]; duplicate { + return fmt.Errorf("claude bridge: user message %d has duplicate tool_result id %q", i+1, id) + } + seen[id] = struct{}{} + } + for id := range expected { + if _, ok := seen[id]; !ok { + return fmt.Errorf("claude bridge: user message %d is missing tool_result for id %q", i+1, id) + } + } + validatedResultMessages[i+1] = struct{}{} + } + for i, msg := range messages { + if _, ok := validatedResultMessages[i]; ok { + continue + } + for _, block := range msg.Content.Blocks { + if block.Type == "tool_result" { + return fmt.Errorf("claude bridge: user message %d has orphan tool_result id %q", i, block.ToolUseID) + } + } + } + return nil +} + +// ============================================================ +// Conversion: Claude Response → OpenAI Response (non-streaming) +// ============================================================ + +// claudeToOpenAIResponseJSON 将 Claude 响应 JSON 转为 OpenAI 兼容的 JSON。 +func claudeToOpenAIResponseJSON(claudeBody []byte) ([]byte, error) { + var cr claudeResponse + if err := json.Unmarshal(claudeBody, &cr); err != nil { + return nil, fmt.Errorf("claude bridge: unmarshal response: %w", err) + } + + if cr.Error != nil { + return nil, fmt.Errorf("claude api error: [%s] %s", cr.Error.Type, cr.Error.Message) + } + + // 构建 OpenAI 格式的 response + oaiResp := map[string]interface{}{ + "id": cr.ID, + "object": "chat.completion", + "model": cr.Model, + "choices": []interface{}{}, + } + + var textContent string + var toolCalls []interface{} + var thinkingBlocks []claudeContentBlock + + for _, block := range cr.Content { + switch block.Type { + case "thinking": + thinkingBlocks = append(thinkingBlocks, block) + case "text": + textContent += block.Text + case "tool_use": + argsStr := string(block.Input) + toolCalls = append(toolCalls, map[string]interface{}{ + "id": block.ID, + "type": "function", + "function": map[string]interface{}{ + "name": block.Name, + "arguments": argsStr, + }, + }) + } + } + + finishReason := claudeStopReasonToOpenAI(cr.StopReason) + message := map[string]interface{}{ + "role": "assistant", + "content": textContent, + } + if len(toolCalls) > 0 { + message["tool_calls"] = toolCalls + } + if len(thinkingBlocks) > 0 { + var parts []string + for _, tb := range thinkingBlocks { + if strings.TrimSpace(tb.Thinking) != "" { + parts = append(parts, tb.Thinking) + } + } + rc := appendClaudeReasoningRoundTrip(strings.Join(parts, "\n\n"), thinkingBlocks) + if rc != "" { + message["reasoning_content"] = rc + } + } + + choice := map[string]interface{}{ + "index": 0, + "message": message, + "finish_reason": finishReason, + } + + oaiResp["choices"] = []interface{}{choice} + + if cr.Usage != nil { + oaiResp["usage"] = map[string]interface{}{ + "prompt_tokens": cr.Usage.InputTokens, + "completion_tokens": cr.Usage.OutputTokens, + "total_tokens": cr.Usage.InputTokens + cr.Usage.OutputTokens, + } + } + + return json.Marshal(oaiResp) +} + +func claudeStopReasonToOpenAI(reason string) string { + switch reason { + case "end_turn": + return "stop" + case "tool_use": + return "tool_calls" + case "max_tokens": + return "length" + case "stop_sequence": + return "stop" + default: + return "stop" + } +} + +// ============================================================ +// Claude HTTP Calls (non-streaming & streaming) +// ============================================================ + +// claudeChatCompletion 执行非流式 Claude API 调用,返回转换后的 OpenAI 格式 JSON。 +func (c *Client) claudeChatCompletion(ctx context.Context, payload interface{}, out interface{}) error { + claudeReq, err := convertOpenAIToClaude(payload) + if err != nil { + return err + } + claudeReq.Stream = false + + body, err := json.Marshal(claudeReq) + if err != nil { + return fmt.Errorf("claude bridge: marshal: %w", err) + } + + baseURL := strings.TrimSuffix(c.config.BaseURL, "/") + if baseURL == "" { + baseURL = "https://api.anthropic.com" + } + + c.logger.Debug("sending Claude chat completion request", + zap.String("model", claudeReq.Model), + zap.Int("payloadSizeKB", len(body)/1024)) + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, baseURL+"/v1/messages", bytes.NewReader(body)) + if err != nil { + return fmt.Errorf("claude bridge: build request: %w", err) + } + c.setClaudeHeaders(req) + + requestStart := time.Now() + resp, err := c.httpClient.Do(req) + if err != nil { + return fmt.Errorf("claude bridge: call api: %w", err) + } + defer resp.Body.Close() + + respBody, err := io.ReadAll(resp.Body) + if err != nil { + return fmt.Errorf("claude bridge: read response: %w", err) + } + + c.logger.Debug("received Claude response", + zap.Int("status", resp.StatusCode), + zap.Duration("duration", time.Since(requestStart)), + zap.Int("responseSizeKB", len(respBody)/1024), + ) + + if resp.StatusCode != http.StatusOK { + c.logger.Warn("Claude chat completion returned non-200", + zap.Int("status", resp.StatusCode), + zap.String("body", string(respBody)), + ) + return &APIError{ + StatusCode: resp.StatusCode, + Body: string(respBody), + } + } + + // 转换为 OpenAI 格式 + oaiJSON, err := claudeToOpenAIResponseJSON(respBody) + if err != nil { + return err + } + + if out != nil { + if err := json.Unmarshal(oaiJSON, out); err != nil { + return fmt.Errorf("claude bridge: unmarshal converted response: %w", err) + } + } + + return nil +} + +// claudeChatCompletionStream 流式调用 Claude API,将 Claude SSE 转换为 OpenAI 兼容的 delta 回调。 +func (c *Client) claudeChatCompletionStream(ctx context.Context, payload interface{}, onDelta func(delta string) error) (string, error) { + claudeReq, err := convertOpenAIToClaude(payload) + if err != nil { + return "", err + } + claudeReq.Stream = true + + body, err := json.Marshal(claudeReq) + if err != nil { + return "", fmt.Errorf("claude bridge: marshal: %w", err) + } + + baseURL := strings.TrimSuffix(c.config.BaseURL, "/") + if baseURL == "" { + baseURL = "https://api.anthropic.com" + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, baseURL+"/v1/messages", bytes.NewReader(body)) + if err != nil { + return "", fmt.Errorf("claude bridge: build request: %w", err) + } + c.setClaudeHeaders(req) + + requestStart := time.Now() + resp, err := c.httpClient.Do(req) + if err != nil { + return "", fmt.Errorf("claude bridge: call api: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + respBody, readErr := io.ReadAll(resp.Body) + if readErr != nil { + return "", fmt.Errorf("claude bridge: read error response: %w", readErr) + } + return "", &APIError{ + StatusCode: resp.StatusCode, + Body: string(respBody), + } + } + + reader := bufio.NewReader(resp.Body) + var full strings.Builder + fullText := "" + + for { + line, readErr := reader.ReadString('\n') + if readErr != nil { + if readErr == io.EOF { + break + } + return full.String(), fmt.Errorf("claude bridge: read stream: %w", readErr) + } + trimmed := strings.TrimSpace(line) + if trimmed == "" || !strings.HasPrefix(trimmed, "data:") { + continue + } + dataStr := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) + if dataStr == "[DONE]" { + break + } + + var event map[string]interface{} + if err := json.Unmarshal([]byte(dataStr), &event); err != nil { + continue + } + + eventType, _ := event["type"].(string) + + switch eventType { + case "content_block_delta": + delta, _ := event["delta"].(map[string]interface{}) + deltaType, _ := delta["type"].(string) + if deltaType == "text_delta" { + text, _ := delta["text"].(string) + if text != "" { + var textOut string + fullText, textOut = normalizeStreamingDelta(fullText, text) + if textOut == "" { + continue + } + full.WriteString(textOut) + if onDelta != nil { + if err := onDelta(textOut); err != nil { + return full.String(), err + } + } + } + } + case "error": + errData, _ := event["error"].(map[string]interface{}) + msg, _ := errData["message"].(string) + return full.String(), fmt.Errorf("claude stream error: %s", msg) + } + } + + c.logger.Debug("received Claude stream completion", + zap.Duration("duration", time.Since(requestStart)), + zap.Int("contentLen", full.Len()), + ) + + return full.String(), nil +} + +// claudeChatCompletionStreamWithToolCalls 流式调用 Claude API,同时处理 content delta 和 tool_calls, +// 返回值与 OpenAI 版本完全一致:(content, toolCalls, finishReason, error)。 +func (c *Client) claudeChatCompletionStreamWithToolCalls( + ctx context.Context, + payload interface{}, + onContentDelta func(delta string) error, +) (string, []StreamToolCall, string, error) { + claudeReq, err := convertOpenAIToClaude(payload) + if err != nil { + return "", nil, "", err + } + claudeReq.Stream = true + + body, err := json.Marshal(claudeReq) + if err != nil { + return "", nil, "", fmt.Errorf("claude bridge: marshal: %w", err) + } + + baseURL := strings.TrimSuffix(c.config.BaseURL, "/") + if baseURL == "" { + baseURL = "https://api.anthropic.com" + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, baseURL+"/v1/messages", bytes.NewReader(body)) + if err != nil { + return "", nil, "", fmt.Errorf("claude bridge: build request: %w", err) + } + c.setClaudeHeaders(req) + + requestStart := time.Now() + resp, err := c.httpClient.Do(req) + if err != nil { + return "", nil, "", fmt.Errorf("claude bridge: call api: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + respBody, readErr := io.ReadAll(resp.Body) + if readErr != nil { + return "", nil, "", fmt.Errorf("claude bridge: read error response: %w", readErr) + } + return "", nil, "", &APIError{ + StatusCode: resp.StatusCode, + Body: string(respBody), + } + } + + reader := bufio.NewReader(resp.Body) + var full strings.Builder + fullText := "" + finishReason := "" + + // 追踪当前正在构建的 content blocks + type toolAccum struct { + id string + name string + args strings.Builder + index int + } + var currentToolCalls []toolAccum + currentBlockIndex := -1 + currentBlockType := "" + + for { + line, readErr := reader.ReadString('\n') + if readErr != nil { + if readErr == io.EOF { + break + } + return full.String(), nil, finishReason, fmt.Errorf("claude bridge: read stream: %w", readErr) + } + trimmed := strings.TrimSpace(line) + if trimmed == "" || !strings.HasPrefix(trimmed, "data:") { + continue + } + dataStr := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) + if dataStr == "[DONE]" { + break + } + + var event map[string]interface{} + if err := json.Unmarshal([]byte(dataStr), &event); err != nil { + continue + } + + eventType, _ := event["type"].(string) + + switch eventType { + case "content_block_start": + idx, _ := event["index"].(float64) + currentBlockIndex = int(idx) + cb, _ := event["content_block"].(map[string]interface{}) + blockType, _ := cb["type"].(string) + currentBlockType = blockType + + if blockType == "tool_use" { + id, _ := cb["id"].(string) + name, _ := cb["name"].(string) + currentToolCalls = append(currentToolCalls, toolAccum{ + id: id, + name: name, + index: currentBlockIndex, + }) + } + + case "content_block_delta": + delta, _ := event["delta"].(map[string]interface{}) + deltaType, _ := delta["type"].(string) + + if deltaType == "text_delta" { + text, _ := delta["text"].(string) + if text != "" { + var textOut string + fullText, textOut = normalizeStreamingDelta(fullText, text) + if textOut == "" { + continue + } + full.WriteString(textOut) + if onContentDelta != nil { + if err := onContentDelta(textOut); err != nil { + return full.String(), nil, finishReason, err + } + } + } + } else if deltaType == "input_json_delta" { + partialJSON, _ := delta["partial_json"].(string) + if partialJSON != "" && currentBlockType == "tool_use" && len(currentToolCalls) > 0 { + currentToolCalls[len(currentToolCalls)-1].args.WriteString(partialJSON) + } + } + + case "content_block_stop": + // block 完成,不需要特殊处理 + + case "message_delta": + delta, _ := event["delta"].(map[string]interface{}) + if sr, ok := delta["stop_reason"].(string); ok { + finishReason = claudeStopReasonToOpenAI(sr) + } + + case "message_stop": + // 消息完成 + + case "error": + errData, _ := event["error"].(map[string]interface{}) + msg, _ := errData["message"].(string) + return full.String(), nil, finishReason, fmt.Errorf("claude stream error: %s", msg) + } + } + + // 转换 tool calls 为 OpenAI 格式的 StreamToolCall + var toolCalls []StreamToolCall + for i, tc := range currentToolCalls { + toolCalls = append(toolCalls, StreamToolCall{ + Index: i, + ID: tc.id, + Type: "function", + FunctionName: tc.name, + FunctionArgsStr: tc.args.String(), + }) + } + + if finishReason == "" { + finishReason = "stop" + } + + c.logger.Debug("received Claude stream completion (tool_calls)", + zap.Duration("duration", time.Since(requestStart)), + zap.Int("contentLen", full.Len()), + zap.Int("toolCalls", len(toolCalls)), + zap.String("finishReason", finishReason), + ) + + return full.String(), toolCalls, finishReason, nil +} + +// ============================================================ +// Helpers +// ============================================================ + +// setClaudeHeaders 设置 Anthropic API 要求的请求头。 +func (c *Client) setClaudeHeaders(req *http.Request) { + req.Header.Set("Content-Type", "application/json") + req.Header.Set("x-api-key", c.config.APIKey) + req.Header.Set("anthropic-version", "2023-06-01") +} + +// isClaude 判断当前配置是否为 Claude provider。 +func (c *Client) isClaude() bool { + return isClaudeProvider(c.config) +} + +func isClaudeProvider(cfg *config.OpenAIConfig) bool { + if cfg == nil { + return false + } + return strings.EqualFold(strings.TrimSpace(cfg.Provider), "claude") || + strings.EqualFold(strings.TrimSpace(cfg.Provider), "anthropic") +} + +// ============================================================ +// Eino HTTP Client Bridge +// ============================================================ + +// NewEinoHTTPClient 为 einoopenai.ChatModelConfig 返回一个 http.Client,包含多层 transport 包装: +// 1. 当 cfg.Provider 为 claude 时,套 claudeRoundTripper,把 OpenAI /chat/completions 透明 +// 桥接为 Anthropic /v1/messages(并把 Claude SSE 翻译回 OpenAI SSE 格式)。 +// 2. reasoningToolChoiceCompatRoundTripper:tool_choice=required/object 时剥离 thinking 字段,避免 +// plan_execute replanner 等强制工具调用与推理模式冲突(部分网关返回 400)。 +// 3. 最外层无条件套 einoSSESanitizingRoundTripper,吞掉中转站发的 SSE 心跳/注释/控制行 +// (": keepalive" / "event: ping" / "retry: 3000" 等),避免 Eino 用的 meguminnnnnnnnn/go-openai +// SDK 在累计超过 300 个非 "data:" 行后抛 "stream has sent too many empty messages"。 +// +// 两层都对调用方完全透明:普通 JSON 响应原样透传,仅当响应 Content-Type 为 text/event-stream 时 +// sanitizer 才会接管 body;data: payload (含 [DONE]、{"error":...}) 一字节不改。 +func NewEinoHTTPClient(cfg *config.OpenAIConfig, base *http.Client) *http.Client { + if base == nil { + base = http.DefaultClient + } + + cloned := *base + transport := base.Transport + if transport == nil { + transport = http.DefaultTransport + } + transport = &reasoningToolChoiceCompatRoundTripper{base: transport, cfg: cfg} + if isClaudeProvider(cfg) { + transport = &claudeRoundTripper{ + base: transport, + config: cfg, + } + } + transport = &einoSSESanitizingRoundTripper{base: transport} + cloned.Transport = transport + return &cloned +} + +// claudeRoundTripper 是一个 http.RoundTripper,用于将 OpenAI 协议透明桥接到 Claude API。 +type claudeRoundTripper struct { + base http.RoundTripper + config *config.OpenAIConfig +} + +func (rt *claudeRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + // 只拦截 chat completions + if !strings.HasSuffix(req.URL.Path, "/chat/completions") { + return rt.base.RoundTrip(req) + } + + // 读取原请求体 + body, err := io.ReadAll(req.Body) + if err != nil { + return nil, fmt.Errorf("claude bridge: read request body: %w", err) + } + _ = req.Body.Close() + + var payload interface{} + if err := json.Unmarshal(body, &payload); err != nil { + return nil, fmt.Errorf("claude bridge: unmarshal request: %w", err) + } + + // 转换为 Claude 请求 + claudeReq, err := convertOpenAIToClaude(payload) + if err != nil { + return nil, err + } + + // 构造 Claude 请求 + baseURL := strings.TrimSuffix(rt.config.BaseURL, "/") + if baseURL == "" { + baseURL = "https://api.anthropic.com" + } + + claudeBody, err := json.Marshal(claudeReq) + if err != nil { + return nil, fmt.Errorf("claude bridge: marshal claude request: %w", err) + } + + newReq, err := http.NewRequestWithContext(req.Context(), http.MethodPost, baseURL+"/v1/messages", bytes.NewReader(claudeBody)) + if err != nil { + return nil, fmt.Errorf("claude bridge: build request: %w", err) + } + newReq.Header.Set("Content-Type", "application/json") + newReq.Header.Set("x-api-key", rt.config.APIKey) + newReq.Header.Set("anthropic-version", "2023-06-01") + + resp, err := rt.base.RoundTrip(newReq) + if err != nil { + return nil, err + } + + // 非 200:尝试把 Claude 错误格式转成 OpenAI 错误格式,便于 Eino 解析 + if resp.StatusCode != http.StatusOK { + bodyBytes, readErr := io.ReadAll(resp.Body) + if readErr != nil { + resp.Body.Close() + return nil, fmt.Errorf("claude bridge: read error response: %w", readErr) + } + resp.Body.Close() + converted := rt.tryConvertClaudeErrorToOpenAI(bodyBytes) + return &http.Response{ + StatusCode: resp.StatusCode, + Header: resp.Header.Clone(), + Body: io.NopCloser(bytes.NewReader(converted)), + ContentLength: int64(len(converted)), + Request: req, + }, nil + } + + // 非流式:一次性转换响应体 + if !claudeReq.Stream { + respBody, readErr := io.ReadAll(resp.Body) + if readErr != nil { + resp.Body.Close() + return nil, fmt.Errorf("claude bridge: read response: %w", readErr) + } + resp.Body.Close() + oaiJSON, err := claudeToOpenAIResponseJSON(respBody) + if err != nil { + return nil, err + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(bytes.NewReader(oaiJSON)), + ContentLength: int64(len(oaiJSON)), + Request: req, + }, nil + } + + // 流式:通过 pipe 实时转换 SSE + pr, pw := io.Pipe() + + // writeLine 将数据写入 pipe,返回 false 表示 pipe 已关闭(消费端断开),应立即退出。 + writeLine := func(data string) bool { + _, err := pw.Write([]byte(data)) + return err == nil + } + + go func() { + defer resp.Body.Close() + + reader := bufio.NewReader(resp.Body) + blockToToolIndex := make(map[int]int) + blockIndexToType := make(map[int]string) + nextToolIndex := 0 + + type thinkingAcc struct { + text strings.Builder + sig strings.Builder + } + thinkingByIndex := make(map[int]*thinkingAcc) + var finishedThinking []claudeContentBlock + + for { + line, readErr := reader.ReadString('\n') + if readErr != nil { + if readErr == io.EOF { + writeLine("data: [DONE]\n\n") + } else { + // 非 EOF 错误:写入错误事件并通知消费端 + oaiErr := map[string]interface{}{ + "error": map[string]interface{}{ + "message": readErr.Error(), + "type": "claude_stream_read_error", + }, + } + b, _ := json.Marshal(oaiErr) + writeLine("data: " + string(b) + "\n\n") + writeLine("data: [DONE]\n\n") + } + pw.Close() + return + } + trimmed := strings.TrimSpace(line) + if trimmed == "" || !strings.HasPrefix(trimmed, "data:") { + continue + } + dataStr := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) + if dataStr == "[DONE]" { + writeLine("data: [DONE]\n\n") + pw.Close() + return + } + + var event map[string]interface{} + if err := json.Unmarshal([]byte(dataStr), &event); err != nil { + continue + } + + eventType, _ := event["type"].(string) + + switch eventType { + case "content_block_start": + blockIdxFlt, _ := event["index"].(float64) + blockIdx := int(blockIdxFlt) + cb, _ := event["content_block"].(map[string]interface{}) + bt, _ := cb["type"].(string) + blockIndexToType[blockIdx] = bt + + if bt == "thinking" { + thinkingByIndex[blockIdx] = &thinkingAcc{} + } + + if bt == "tool_use" { + id, _ := cb["id"].(string) + name, _ := cb["name"].(string) + blockToToolIndex[blockIdx] = nextToolIndex + toolIdx := nextToolIndex + nextToolIndex++ + + oaiChunk := map[string]interface{}{ + "choices": []map[string]interface{}{ + { + "delta": map[string]interface{}{ + "tool_calls": []map[string]interface{}{ + { + "index": toolIdx, + "id": id, + "type": "function", + "function": map[string]interface{}{ + "name": name, + }, + }, + }, + }, + }, + }, + } + b, _ := json.Marshal(oaiChunk) + if !writeLine("data: " + string(b) + "\n\n") { + pw.Close() + return + } + } + + case "content_block_delta": + blockIdxFlt, _ := event["index"].(float64) + blockIdx := int(blockIdxFlt) + delta, _ := event["delta"].(map[string]interface{}) + dt, _ := delta["type"].(string) + + if dt == "thinking_delta" { + tPart, _ := delta["thinking"].(string) + if tPart != "" { + if acc := thinkingByIndex[blockIdx]; acc != nil { + acc.text.WriteString(tPart) + } + oaiChunk := map[string]interface{}{ + "choices": []map[string]interface{}{ + { + "delta": map[string]interface{}{ + "reasoning_content": tPart, + }, + }, + }, + } + b, _ := json.Marshal(oaiChunk) + if !writeLine("data: " + string(b) + "\n\n") { + pw.Close() + return + } + } + } else if dt == "signature_delta" { + sigPart, _ := delta["signature"].(string) + if sigPart != "" { + if acc := thinkingByIndex[blockIdx]; acc != nil { + acc.sig.WriteString(sigPart) + } + } + } else if dt == "text_delta" { + text, _ := delta["text"].(string) + oaiChunk := map[string]interface{}{ + "choices": []map[string]interface{}{ + { + "delta": map[string]interface{}{ + "content": text, + }, + }, + }, + } + b, _ := json.Marshal(oaiChunk) + if !writeLine("data: " + string(b) + "\n\n") { + pw.Close() + return + } + } else if dt == "input_json_delta" { + partial, _ := delta["partial_json"].(string) + if partial != "" { + if toolIdx, ok := blockToToolIndex[blockIdx]; ok { + oaiChunk := map[string]interface{}{ + "choices": []map[string]interface{}{ + { + "delta": map[string]interface{}{ + "tool_calls": []map[string]interface{}{ + { + "index": toolIdx, + "function": map[string]interface{}{ + "arguments": partial, + }, + }, + }, + }, + }, + }, + } + b, _ := json.Marshal(oaiChunk) + if !writeLine("data: " + string(b) + "\n\n") { + pw.Close() + return + } + } + } + } + + case "content_block_stop": + blockIdxFlt, _ := event["index"].(float64) + blockIdx := int(blockIdxFlt) + bt := blockIndexToType[blockIdx] + if bt == "thinking" { + if acc := thinkingByIndex[blockIdx]; acc != nil { + finishedThinking = append(finishedThinking, claudeContentBlock{ + Type: "thinking", + Thinking: acc.text.String(), + Signature: acc.sig.String(), + }) + delete(thinkingByIndex, blockIdx) + } + } + + case "message_delta": + d, _ := event["delta"].(map[string]interface{}) + if sr, ok := d["stop_reason"].(string); ok { + finishReason := claudeStopReasonToOpenAI(sr) + oaiChunk := map[string]interface{}{ + "choices": []map[string]interface{}{ + { + "delta": map[string]interface{}{}, + "finish_reason": finishReason, + }, + }, + } + b, _ := json.Marshal(oaiChunk) + if !writeLine("data: " + string(b) + "\n\n") { + pw.Close() + return + } + } + + case "message_stop": + if len(finishedThinking) > 0 { + suffix := appendClaudeReasoningRoundTrip("", finishedThinking) + if strings.TrimSpace(suffix) != "" { + oaiChunk := map[string]interface{}{ + "choices": []map[string]interface{}{ + { + "delta": map[string]interface{}{ + "reasoning_content": suffix, + }, + }, + }, + } + b, _ := json.Marshal(oaiChunk) + if !writeLine("data: " + string(b) + "\n\n") { + pw.Close() + return + } + } + } + writeLine("data: [DONE]\n\n") + pw.Close() + return + + case "error": + errData, _ := event["error"].(map[string]interface{}) + msg, _ := errData["message"].(string) + oaiChunk := map[string]interface{}{ + "error": map[string]interface{}{ + "message": msg, + "type": "claude_stream_error", + }, + } + b, _ := json.Marshal(oaiChunk) + writeLine("data: " + string(b) + "\n\n") + writeLine("data: [DONE]\n\n") + pw.Close() + return + } + } + }() + + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + }, + Body: pr, + Request: req, + }, nil +} + +// tryConvertClaudeErrorToOpenAI 尝试把 Claude 错误格式转换为 OpenAI 错误格式 JSON。 +func (rt *claudeRoundTripper) tryConvertClaudeErrorToOpenAI(body []byte) []byte { + var ce struct { + Type string `json:"type"` + Error struct { + Type string `json:"type"` + Message string `json:"message"` + } `json:"error"` + } + if err := json.Unmarshal(body, &ce); err != nil || ce.Error.Message == "" { + return body + } + oaiErr := map[string]interface{}{ + "error": map[string]interface{}{ + "message": ce.Error.Message, + "type": ce.Error.Type, + "code": ce.Type, + }, + } + b, _ := json.Marshal(oaiErr) + return b +} diff --git a/internal/openai/claude_reasoning_roundtrip.go b/internal/openai/claude_reasoning_roundtrip.go new file mode 100644 index 00000000..1eae4c67 --- /dev/null +++ b/internal/openai/claude_reasoning_roundtrip.go @@ -0,0 +1,81 @@ +package openai + +import ( + "encoding/json" + "strings" +) + +// claudeReasoningRoundTripSep separates human-readable reasoning from a JSON payload of +// Anthropic thinking blocks (with signatures) for multi-turn extended thinking + tools. +// Not shown in UI (see DisplayReasoningContent). +const claudeReasoningRoundTripSep = "\n---CSAI_CLAUDE_THINKING_BLOCKS---\n" + +// DisplayReasoningContent returns reasoning text suitable for the UI (strips internal +// Claude round-trip JSON suffix). Safe for DeepSeek/plain reasoning strings (no-op). +func DisplayReasoningContent(s string) string { + s = strings.TrimSpace(s) + if s == "" { + return "" + } + i := strings.LastIndex(s, claudeReasoningRoundTripSep) + if i < 0 { + return s + } + return strings.TrimSpace(s[:i]) +} + +func appendClaudeReasoningRoundTrip(display string, blocks []claudeContentBlock) string { + var payload []map[string]string + for _, b := range blocks { + if b.Type != "thinking" { + continue + } + payload = append(payload, map[string]string{ + "type": b.Type, + "thinking": b.Thinking, + "signature": b.Signature, + }) + } + if len(payload) == 0 { + return strings.TrimSpace(display) + } + js, err := json.Marshal(payload) + if err != nil { + return strings.TrimSpace(display) + } + d := strings.TrimSpace(display) + if d == "" { + return claudeReasoningRoundTripSep + string(js) + } + return d + claudeReasoningRoundTripSep + string(js) +} + +// parseClaudeReasoningAssistantBlocks extracts Anthropic thinking blocks from an OpenAI-style +// reasoning_content string. When no suffix is present, blocks is nil (caller must not invent signatures). +func parseClaudeReasoningAssistantBlocks(reasoningContent string) (display string, blocks []claudeContentBlock) { + reasoningContent = strings.TrimSpace(reasoningContent) + if reasoningContent == "" { + return "", nil + } + idx := strings.LastIndex(reasoningContent, claudeReasoningRoundTripSep) + if idx < 0 { + return reasoningContent, nil + } + display = strings.TrimSpace(reasoningContent[:idx]) + jsonPart := strings.TrimSpace(reasoningContent[idx+len(claudeReasoningRoundTripSep):]) + var arr []struct { + Type string `json:"type"` + Thinking string `json:"thinking"` + Signature string `json:"signature"` + } + if err := json.Unmarshal([]byte(jsonPart), &arr); err != nil { + return reasoningContent, nil + } + for _, x := range arr { + if x.Type != "thinking" { + continue + } + blocks = append(blocks, claudeContentBlock{Type: "thinking", Thinking: x.Thinking, Signature: x.Signature}) + } + return display, blocks +} diff --git a/internal/openai/claude_reasoning_roundtrip_test.go b/internal/openai/claude_reasoning_roundtrip_test.go new file mode 100644 index 00000000..67dc25ac --- /dev/null +++ b/internal/openai/claude_reasoning_roundtrip_test.go @@ -0,0 +1,211 @@ +package openai + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestDisplayReasoningContent(t *testing.T) { + raw := "hello" + claudeReasoningRoundTripSep + `[{"type":"thinking","thinking":"x","signature":"sig"}]` + if d := DisplayReasoningContent(raw); d != "hello" { + t.Fatalf("got %q", d) + } + if DisplayReasoningContent("plain") != "plain" { + t.Fatal() + } +} + +func TestAppendParseClaudeReasoningRoundTrip(t *testing.T) { + blocks := []claudeContentBlock{ + {Type: "thinking", Thinking: "a", Signature: "s1"}, + {Type: "thinking", Thinking: "b", Signature: "s2"}, + } + s := appendClaudeReasoningRoundTrip("sum", blocks) + if !strings.Contains(s, claudeReasoningRoundTripSep) { + t.Fatal("missing sep") + } + display, back := parseClaudeReasoningAssistantBlocks(s) + if display != "sum" || len(back) != 2 { + t.Fatalf("display=%q len=%d", display, len(back)) + } + if back[0].Signature != "s1" || back[1].Thinking != "b" { + t.Fatalf("%+v", back) + } +} + +func TestConvertOpenAIToClaude_AssistantReasoningReplay(t *testing.T) { + rc := appendClaudeReasoningRoundTrip("vis", []claudeContentBlock{ + {Type: "thinking", Thinking: "t1", Signature: "sig1"}, + }) + payload := map[string]interface{}{ + "model": "claude-3-5-sonnet-latest", + "messages": []interface{}{ + map[string]interface{}{ + "role": "assistant", + "content": "out", + "reasoning_content": rc, + }, + }, + } + req, err := convertOpenAIToClaude(payload) + if err != nil { + t.Fatal(err) + } + if len(req.Messages) != 1 { + t.Fatalf("messages=%d", len(req.Messages)) + } + blocks := req.Messages[0].Content.Blocks + if len(blocks) < 2 { + t.Fatalf("blocks=%d", len(blocks)) + } + if blocks[0].Type != "thinking" || blocks[0].Signature != "sig1" { + t.Fatalf("first block %+v", blocks[0]) + } + foundText := false + for _, b := range blocks { + if b.Type == "text" && b.Text == "out" { + foundText = true + } + } + if !foundText { + t.Fatalf("blocks=%+v", blocks) + } +} + +func TestConvertOpenAIToClaude_OutputConfigEffort(t *testing.T) { + payload := map[string]interface{}{ + "model": "claude-opus-4-8", + "messages": []interface{}{ + map[string]interface{}{"role": "user", "content": "hi"}, + }, + "thinking": map[string]interface{}{ + "type": "adaptive", + "display": "summarized", + }, + "output_config": map[string]interface{}{ + "effort": "high", + }, + } + req, err := convertOpenAIToClaude(payload) + if err != nil { + t.Fatal(err) + } + if len(req.Thinking) == 0 { + t.Fatal("expected thinking") + } + if len(req.OutputConfig) == 0 { + t.Fatal("expected output_config") + } + var oc map[string]interface{} + if err := json.Unmarshal(req.OutputConfig, &oc); err != nil { + t.Fatal(err) + } + if oc["effort"] != "high" { + t.Fatalf("effort=%v", oc["effort"]) + } +} + +func TestConvertOpenAIToClaudeAcceptsCompleteMultiToolBatch(t *testing.T) { + payload := map[string]interface{}{ + "model": "claude-test", + "messages": []interface{}{ + map[string]interface{}{ + "role": "assistant", + "tool_calls": []interface{}{ + map[string]interface{}{"id": "c1", "function": map[string]interface{}{"name": "one", "arguments": "{}"}}, + map[string]interface{}{"id": "c2", "function": map[string]interface{}{"name": "two", "arguments": "{}"}}, + }, + }, + map[string]interface{}{"role": "tool", "tool_call_id": "c1", "content": "r1"}, + map[string]interface{}{"role": "tool", "tool_call_id": "c2", "content": "r2"}, + map[string]interface{}{"role": "assistant", "content": "done"}, + }, + } + req, err := convertOpenAIToClaude(payload) + if err != nil { + t.Fatal(err) + } + if len(req.Messages) != 3 || len(req.Messages[1].Content.Blocks) != 2 { + t.Fatalf("unexpected converted messages: %+v", req.Messages) + } +} + +func TestConvertOpenAIToClaudeRejectsMissingToolResult(t *testing.T) { + payload := map[string]interface{}{ + "model": "claude-test", + "messages": []interface{}{ + map[string]interface{}{ + "role": "assistant", + "tool_calls": []interface{}{ + map[string]interface{}{"id": "c1", "function": map[string]interface{}{"name": "one", "arguments": "{}"}}, + map[string]interface{}{"id": "c2", "function": map[string]interface{}{"name": "two", "arguments": "{}"}}, + }, + }, + map[string]interface{}{"role": "tool", "tool_call_id": "c1", "content": "r1"}, + }, + } + _, err := convertOpenAIToClaude(payload) + if err == nil || !strings.Contains(err.Error(), "missing tool_result") { + t.Fatalf("expected missing tool_result error, got %v", err) + } +} + +func TestConvertOpenAIToClaudeRejectsOrphanToolResult(t *testing.T) { + payload := map[string]interface{}{ + "model": "claude-test", + "messages": []interface{}{ + map[string]interface{}{"role": "user", "content": "start"}, + map[string]interface{}{"role": "tool", "tool_call_id": "orphan", "content": "result"}, + }, + } + _, err := convertOpenAIToClaude(payload) + if err == nil || !strings.Contains(err.Error(), "orphan tool_result") { + t.Fatalf("expected orphan tool_result error, got %v", err) + } +} + +func TestClaudeToOpenAIResponseJSON_Thinking(t *testing.T) { + claudeBody := []byte(`{ + "id":"msg_1","type":"message","role":"assistant","model":"x","stop_reason":"end_turn", + "content":[ + {"type":"thinking","thinking":"step","signature":"sigx"}, + {"type":"text","text":"hi"} + ] + }`) + oai, err := claudeToOpenAIResponseJSON(claudeBody) + if err != nil { + t.Fatal(err) + } + var wrap map[string]interface{} + if err := json.Unmarshal(oai, &wrap); err != nil { + t.Fatal(err) + } + choices := wrap["choices"].([]interface{}) + ch0 := choices[0].(map[string]interface{}) + msg := ch0["message"].(map[string]interface{}) + rc, _ := msg["reasoning_content"].(string) + if !strings.Contains(rc, "step") || !strings.Contains(rc, claudeReasoningRoundTripSep) { + t.Fatalf("reasoning_content=%q", rc) + } + if msg["content"] != "hi" { + t.Fatal() + } +} + +func TestConvertOpenAIToClaudeMapsMaxCompletionTokens(t *testing.T) { + req, err := convertOpenAIToClaude(map[string]interface{}{ + "model": "claude-test", + "max_completion_tokens": float64(16384), + "max_tokens": float64(1024), + "messages": []interface{}{ + map[string]interface{}{"role": "user", "content": "hello"}, + }, + }) + if err != nil { + t.Fatal(err) + } + if req.MaxTokens != 16384 { + t.Fatalf("max tokens=%d, want 16384", req.MaxTokens) + } +} diff --git a/internal/openai/eino_sse_sanitizer.go b/internal/openai/eino_sse_sanitizer.go new file mode 100644 index 00000000..43e07d5b --- /dev/null +++ b/internal/openai/eino_sse_sanitizer.go @@ -0,0 +1,149 @@ +package openai + +// eino_sse_sanitizer.go 解决 Eino 走 meguminnnnnnnnn/go-openai SDK 时, +// 中转站心跳/SSE 控制行累计 > 300 行触发 ErrTooManyEmptyStreamMessages +// (报错文案: "stream has sent too many empty messages")的问题。 +// +// 触发链路: +// einoopenai.NewChatModel +// → eino-ext/libs/acl/openai → meguminnnnnnnnn/go-openai +// → streamReader.processLines() 对所有非 "data:" 行计数, > 300 即抛错。 +// +// 中转站常见的非 data: 行(合法 SSE 但 SDK 不接受): +// ":" / ": keepalive" / ": ping" / "event: ping" / "retry: 3000" +// 以及思考型模型 prefill 期间穿插的大量心跳。 +// +// 兜底策略: 在 HTTP transport 层把响应 Body 包一层 reader, 只放行 "data:" +// 开头的行, 把心跳/注释/事件类型行就地吞掉。下游 SDK 永远见不到非 data: 行, +// 计数器始终为 0, 该错误不可能再发生。 +// +// 该层对调用方完全透明: +// - 仅当响应 Content-Type 是 text/event-stream 时介入;普通 JSON 响应原样透传 +// - data: payload (含 [DONE] 与 {"error":...}) 一字节不改 +// - 上游真断流 (EOF / connection reset / context cancel) 原样透传 + +import ( + "bufio" + "bytes" + "io" + "net/http" + "strings" +) + +const ( + // einoSSEReaderBufSize 给 bufio 一个较大的初始缓冲, 避免单行大 JSON chunk + // (含工具调用 arguments / reasoning_content) 频繁触发缓冲区扩容。 + einoSSEReaderBufSize = 64 * 1024 +) + +// einoSSESanitizingRoundTripper 包装下游 RoundTripper, 对 SSE 响应做行级清洗。 +type einoSSESanitizingRoundTripper struct { + base http.RoundTripper +} + +func (rt *einoSSESanitizingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + resp, err := rt.base.RoundTrip(req) + if err != nil || resp == nil { + return resp, err + } + if !isSSEResponse(resp) { + return resp, nil + } + resp.Body = newEinoSSESanitizingBody(resp.Body) + return resp, nil +} + +// isSSEResponse 仅对 200 + text/event-stream 的响应做清洗; +// 错误响应 (4xx/5xx 通常是 application/json) 不动, 由 SDK 走原错误路径。 +func isSSEResponse(resp *http.Response) bool { + if resp.StatusCode != http.StatusOK { + return false + } + ct := resp.Header.Get("Content-Type") + if ct == "" { + return false + } + ct = strings.ToLower(strings.TrimSpace(ct)) + // 兼容 "text/event-stream", "text/event-stream; charset=utf-8" 等。 + return strings.HasPrefix(ct, "text/event-stream") +} + +// einoSSESanitizingBody 是包装后的响应体: 只放行 data: 行, 其它行吞掉。 +type einoSSESanitizingBody struct { + upstream io.ReadCloser + reader *bufio.Reader + pending []byte // 已清洗、待返回给下游的字节 (永远以 \n 结尾的完整 data: 行) + err error // upstream 终态错误 (io.EOF 或网络错误) +} + +func newEinoSSESanitizingBody(body io.ReadCloser) *einoSSESanitizingBody { + return &einoSSESanitizingBody{ + upstream: body, + reader: bufio.NewReaderSize(body, einoSSEReaderBufSize), + } +} + +func (b *einoSSESanitizingBody) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + if len(b.pending) > 0 { + n := copy(p, b.pending) + b.pending = b.pending[n:] + return n, nil + } + + // 从上游读, 直到攒出一行 data: 或拿到终态。 + // 单次循环可能丢弃任意多行心跳, 但只放行至多一行 data: 后退出, + // 避免一次 Read 阻塞过久 / pending 缓冲过大。 + for b.err == nil { + line, err := b.reader.ReadBytes('\n') + if len(line) > 0 { + if isPassThroughSSELine(line) { + if line[len(line)-1] != '\n' { + line = append(line, '\n') + } + b.pending = line + if err != nil { + b.err = err + } + break + } + // 非 data: 行 (空行 / ":" 注释 / event: / retry: / id: / 任何裸文本) + // 全部吞掉, 不向下游透出, 继续循环读下一行。 + } + if err != nil { + b.err = err + break + } + } + + if len(b.pending) > 0 { + n := copy(p, b.pending) + b.pending = b.pending[n:] + return n, nil + } + return 0, b.err +} + +func (b *einoSSESanitizingBody) Close() error { + return b.upstream.Close() +} + +// isPassThroughSSELine 判定该行是否需要原样放行给下游 SDK。 +// 仅 "data:" (大小写不敏感, 可有任意前导空白) 开头的行需要保留。 +// 注意: 不能用 TrimSpace 去尾部换行后再判, 否则 " data: x" 会被误判; +// 我们只 trim 前导空白, 与 SDK 内部 TrimSpace 后再正则 ^data:\s* 的语义一致。 +func isPassThroughSSELine(line []byte) bool { + trimmed := bytes.TrimLeft(line, " \t") + if len(trimmed) < 5 { + return false + } + // 大小写不敏感比较前 5 字节是否为 "data:"。SSE 规范要求字段名小写, + // 但宽松匹配可以兼容个别中转站的非规范实现。 + return (trimmed[0] == 'd' || trimmed[0] == 'D') && + (trimmed[1] == 'a' || trimmed[1] == 'A') && + (trimmed[2] == 't' || trimmed[2] == 'T') && + (trimmed[3] == 'a' || trimmed[3] == 'A') && + trimmed[4] == ':' +} diff --git a/internal/openai/eino_sse_sanitizer_test.go b/internal/openai/eino_sse_sanitizer_test.go new file mode 100644 index 00000000..ef52db39 --- /dev/null +++ b/internal/openai/eino_sse_sanitizer_test.go @@ -0,0 +1,303 @@ +package openai + +import ( + "bufio" + "bytes" + "errors" + "io" + "net/http" + "net/http/httptest" + "regexp" + "strings" + "testing" +) + +// 复现 meguminnnnnnnnn/go-openai 的 SSE 行计数算法 (默认 limit=300): +// - 逐行读 +// - 非 "data:" 行 (空行 / ":" 注释 / event: / retry:) 累计 emptyMessagesCount +// - > 300 抛 ErrTooManyEmptyStreamMessages +// - 遇到 data: 行 reset, 返回 payload +// +// 这一算法与上游 SDK 的 stream_reader.go processLines() 严格一致 (验证依据见 +// /Users/temp/go/pkg/mod/github.com/meguminnnnnnnnn/go-openai@v0.1.2/stream_reader.go)。 +// 测试中只复刻 "限制触发" 这一行为, 用来回归验证 sanitizer 的根因修复。 +var errTooManyEmptyStreamMessages = errors.New("stream has sent too many empty messages") + +func sdkLikeRecvAll(body io.Reader, limit uint) ([]string, error) { + headerData := regexp.MustCompile(`^data:\s*`) + r := bufio.NewReader(body) + var payloads []string + for { + var emptyMessagesCount uint + var payload []byte + for { + line, err := r.ReadBytes('\n') + if err != nil { + if err == io.EOF { + return payloads, nil + } + return payloads, err + } + noSpace := bytes.TrimSpace(line) + if !headerData.Match(noSpace) { + emptyMessagesCount++ + if emptyMessagesCount > limit { + return payloads, errTooManyEmptyStreamMessages + } + continue + } + payload = headerData.ReplaceAll(noSpace, nil) + break + } + if string(payload) == "[DONE]" { + return payloads, nil + } + payloads = append(payloads, string(payload)) + } +} + +func newSSEServer(t *testing.T, body string, contentType string, status int) *httptest.Server { + t.Helper() + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + if contentType != "" { + w.Header().Set("Content-Type", contentType) + } + w.WriteHeader(status) + _, _ = io.WriteString(w, body) + })) +} + +func sanitizingClient(base *http.Client) *http.Client { + if base == nil { + base = &http.Client{} + } + cloned := *base + transport := base.Transport + if transport == nil { + transport = http.DefaultTransport + } + cloned.Transport = &einoSSESanitizingRoundTripper{base: transport} + return &cloned +} + +func readAll(t *testing.T, body io.ReadCloser) string { + t.Helper() + defer body.Close() + out, err := io.ReadAll(body) + if err != nil { + t.Fatalf("read body: %v", err) + } + return string(out) +} + +// 1) 仅 data: 行 → 一字节不改地透传。 +func TestSSESanitizer_PassesDataLinesUnchanged(t *testing.T) { + body := "data: {\"a\":1}\ndata: {\"b\":2}\ndata: [DONE]\n" + srv := newSSEServer(t, body, "text/event-stream", 200) + defer srv.Close() + + resp, err := sanitizingClient(nil).Get(srv.URL) + if err != nil { + t.Fatalf("get: %v", err) + } + got := readAll(t, resp.Body) + if got != body { + t.Fatalf("body mismatch:\nwant %q\ngot %q", body, got) + } +} + +// 2) 心跳/注释/事件类型行被吞掉, 仅保留 data: 行。 +func TestSSESanitizer_DropsHeartbeatsAndControlLines(t *testing.T) { + body := strings.Join([]string{ + ": keepalive", + "", + "event: ping", + "retry: 3000", + "id: 42", + "data: {\"x\":1}", + ": ping", + "", + "data: {\"x\":2}", + "data: [DONE]", + "", + }, "\n") + srv := newSSEServer(t, body, "text/event-stream", 200) + defer srv.Close() + + resp, err := sanitizingClient(nil).Get(srv.URL) + if err != nil { + t.Fatalf("get: %v", err) + } + got := readAll(t, resp.Body) + want := "data: {\"x\":1}\ndata: {\"x\":2}\ndata: [DONE]\n" + if got != want { + t.Fatalf("sanitized body mismatch:\nwant %q\ngot %q", want, got) + } +} + +// 3) 根因回归: 上游堆 500 行心跳后才发 data:, 原始 SDK 算法会抛 +// ErrTooManyEmptyStreamMessages, sanitize 之后必须能正常拿到所有 data:。 +func TestSSESanitizer_ProtectsAgainstTooManyEmptyMessages(t *testing.T) { + const heartbeats = 500 + var buf bytes.Buffer + for i := 0; i < heartbeats; i++ { + buf.WriteString(": keepalive\n") + } + buf.WriteString("data: {\"chunk\":1}\n") + buf.WriteString("data: {\"chunk\":2}\n") + buf.WriteString("data: [DONE]\n") + + t.Run("baseline_without_sanitizer_must_fail", func(t *testing.T) { + _, err := sdkLikeRecvAll(bytes.NewReader(buf.Bytes()), 300) + if !errors.Is(err, errTooManyEmptyStreamMessages) { + t.Fatalf("expected ErrTooManyEmptyStreamMessages, got %v", err) + } + }) + + t.Run("with_sanitizer_must_succeed", func(t *testing.T) { + srv := newSSEServer(t, buf.String(), "text/event-stream", 200) + defer srv.Close() + + resp, err := sanitizingClient(nil).Get(srv.URL) + if err != nil { + t.Fatalf("get: %v", err) + } + defer resp.Body.Close() + + payloads, err := sdkLikeRecvAll(resp.Body, 300) + if err != nil { + t.Fatalf("sdk-like recv after sanitize: %v", err) + } + want := []string{`{"chunk":1}`, `{"chunk":2}`} + if len(payloads) != len(want) { + t.Fatalf("payload count mismatch: want %d got %d (%v)", len(want), len(payloads), payloads) + } + for i, w := range want { + if payloads[i] != w { + t.Fatalf("payload[%d] mismatch: want %q got %q", i, w, payloads[i]) + } + } + }) +} + +// 4) 心跳穿插在 data: 之间也能正确清洗 (思考型模型 prefill 期间常见)。 +func TestSSESanitizer_HeartbeatsInterleavedWithData(t *testing.T) { + var buf bytes.Buffer + buf.WriteString("data: {\"chunk\":1}\n") + for i := 0; i < 400; i++ { + buf.WriteString(": keepalive\n") + } + buf.WriteString("data: {\"chunk\":2}\n") + buf.WriteString("data: [DONE]\n") + + srv := newSSEServer(t, buf.String(), "text/event-stream", 200) + defer srv.Close() + + resp, err := sanitizingClient(nil).Get(srv.URL) + if err != nil { + t.Fatalf("get: %v", err) + } + defer resp.Body.Close() + + payloads, err := sdkLikeRecvAll(resp.Body, 300) + if err != nil { + t.Fatalf("sdk-like recv: %v", err) + } + if got, want := len(payloads), 2; got != want { + t.Fatalf("payload count: want %d got %d", want, got) + } +} + +// 5) 非 SSE 响应 (例如非流式 JSON) 不应被 sanitizer 介入。 +func TestSSESanitizer_PassesNonSSEResponseUntouched(t *testing.T) { + body := `{"id":"x","object":"chat.completion","choices":[]}` + srv := newSSEServer(t, body, "application/json", 200) + defer srv.Close() + + resp, err := sanitizingClient(nil).Get(srv.URL) + if err != nil { + t.Fatalf("get: %v", err) + } + got := readAll(t, resp.Body) + if got != body { + t.Fatalf("non-SSE body must be untouched:\nwant %q\ngot %q", body, got) + } +} + +// 6) 错误响应 (4xx/5xx) 不应被 sanitize, 即使 Content-Type 是 SSE 也不动, +// 避免吞掉类似 "data: " 之外的错误正文。 +func TestSSESanitizer_PassesNon200Untouched(t *testing.T) { + body := `{"error":{"message":"rate limit"}}` + srv := newSSEServer(t, body, "text/event-stream", 429) + defer srv.Close() + + resp, err := sanitizingClient(nil).Get(srv.URL) + if err != nil { + t.Fatalf("get: %v", err) + } + got := readAll(t, resp.Body) + if got != body { + t.Fatalf("error body must be untouched:\nwant %q\ngot %q", body, got) + } +} + +// 7) data: 行末尾若缺 \n (异常上游) sanitizer 也补齐, 保证下游按行解析。 +func TestSSESanitizer_AppendsTrailingNewlineIfMissing(t *testing.T) { + body := "data: {\"a\":1}" + srv := newSSEServer(t, body, "text/event-stream", 200) + defer srv.Close() + + resp, err := sanitizingClient(nil).Get(srv.URL) + if err != nil { + t.Fatalf("get: %v", err) + } + got := readAll(t, resp.Body) + want := "data: {\"a\":1}\n" + if got != want { + t.Fatalf("trailing newline:\nwant %q\ngot %q", want, got) + } +} + +// 8) 大 chunk (一行数十 KB) 也能完整透传, 不被切断。 +func TestSSESanitizer_LargeDataLinePassesIntact(t *testing.T) { + huge := strings.Repeat("x", 80*1024) + body := "data: {\"big\":\"" + huge + "\"}\ndata: [DONE]\n" + srv := newSSEServer(t, body, "text/event-stream", 200) + defer srv.Close() + + resp, err := sanitizingClient(nil).Get(srv.URL) + if err != nil { + t.Fatalf("get: %v", err) + } + got := readAll(t, resp.Body) + if got != body { + t.Fatalf("large body length mismatch: want %d got %d", len(body), len(got)) + } +} + +// 9) isPassThroughSSELine 单元覆盖。 +func TestIsPassThroughSSELine(t *testing.T) { + cases := []struct { + line string + want bool + }{ + {"data: {\"a\":1}\n", true}, + {"DATA: x\n", true}, + {" data: x\n", true}, + {"data:\n", true}, + {"\n", false}, + {"\r\n", false}, + {": keepalive\n", false}, + {":\n", false}, + {"event: ping\n", false}, + {"retry: 3000\n", false}, + {"id: 42\n", false}, + {"datax: y\n", false}, + {"da", false}, + } + for _, c := range cases { + if got := isPassThroughSSELine([]byte(c.line)); got != c.want { + t.Errorf("isPassThroughSSELine(%q) = %v, want %v", c.line, got, c.want) + } + } +} diff --git a/internal/openai/normalize_streaming_delta_test.go b/internal/openai/normalize_streaming_delta_test.go new file mode 100644 index 00000000..6959b590 --- /dev/null +++ b/internal/openai/normalize_streaming_delta_test.go @@ -0,0 +1,56 @@ +package openai + +import "testing" + +func TestNormalizeStreamingDelta_RepeatedCharBoundary(t *testing.T) { + // 流式在重复数字边界分片:不得把 "43" 的首字符与 "194" 尾字符误合并。 + cur, d := normalizeStreamingDelta("https://x:194", "43") + if want := "https://x:19443"; cur != want { + t.Fatalf("next: want %q got %q", want, cur) + } + if d != "43" { + t.Fatalf("delta: want %q got %q", "43", d) + } +} + +func TestNormalizeStreamingDelta_CumulativePrefix(t *testing.T) { + cur, d := normalizeStreamingDelta("今天", "今天天气") + if cur != "今天天气" || d != "天气" { + t.Fatalf("got cur=%q d=%q", cur, d) + } +} + +func TestNormalizeStreamingDelta_FullRetransmit(t *testing.T) { + cur, d := normalizeStreamingDelta("今天", "今天") + if d != "" || cur != "今天" { + t.Fatalf("got cur=%q d=%q", cur, d) + } +} + +func TestNormalizeStreamingDelta_SingleRuneRepeated(t *testing.T) { + cur, d := normalizeStreamingDelta("呀", "呀") + if want := "呀呀"; cur != want { + t.Fatalf("next: want %q got %q", want, cur) + } + if d != "呀" { + t.Fatalf("delta: want %q got %q", "呀", d) + } + cur, d = normalizeStreamingDelta("4", "4") + if want := "44"; cur != want { + t.Fatalf("next: want %q got %q", want, cur) + } + if d != "4" { + t.Fatalf("delta: want %q got %q", "4", d) + } +} + +func TestNormalizeStreamingDelta_CumulativeExtendsNumber(t *testing.T) { + // 已缓冲 "194" 后收到累计串 "19443"(注意 "1943" 并非 "19443" 的前缀,不能靠误写的中间态测 HasPrefix)。 + cur, d := normalizeStreamingDelta("194", "19443") + if want := "19443"; cur != want { + t.Fatalf("next: want %q got %q", want, cur) + } + if d != "43" { + t.Fatalf("delta: want %q got %q", "43", d) + } +} diff --git a/internal/openai/openai.go b/internal/openai/openai.go new file mode 100644 index 00000000..b4b36515 --- /dev/null +++ b/internal/openai/openai.go @@ -0,0 +1,616 @@ +package openai + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "sort" + "strings" + "time" + "unicode/utf8" + + "cyberstrike-ai/internal/config" + + "go.uber.org/zap" +) + +// Client 统一封装与OpenAI兼容模型交互的HTTP客户端。 +type Client struct { + httpClient *http.Client + config *config.OpenAIConfig + logger *zap.Logger +} + +// APIError 表示OpenAI接口返回的非200错误。 +type APIError struct { + StatusCode int + Body string +} + +func (e *APIError) Error() string { + return fmt.Sprintf("openai api error: status=%d body=%s", e.StatusCode, e.Body) +} + +// normalizeStreamingDelta 将可能是“累计片段/重发片段”的内容归一化为“纯增量”。 +// 部分兼容网关会返回累计 content;若直接 append 会出现重复文本。 +// +// 注意: +// - 不做「任意后缀与前缀重叠」合并;流式可能在重复字符边界分片("194"+"43"→"19443")。 +// - HasPrefix 仅在 incoming 严格长于 current 时视为累计全文,否则会把分片产生的第二个相同 +// 单字/单码点(叠字、44、22 等)误判为「整段重复」而吞字。 +// - incoming==current 仅当 current 长度 >1 个码点时才视为整包重发;单码点重复必须走拼接。 +// - 不再使用「current 以 incoming 结尾则丢弃」:否则 "1943"+"43" 会误吞增量(19443 显示成 1943)。 +// 若网关重复发送尾部片段,应重复送完整累计串,由 HasPrefix 分支去重。 +func normalizeStreamingDelta(current, incoming string) (next, delta string) { + if incoming == "" { + return current, "" + } + if current == "" { + return incoming, incoming + } + if strings.HasPrefix(incoming, current) && len(incoming) > len(current) { + return incoming, incoming[len(current):] + } + if incoming == current && utf8.RuneCountInString(current) > 1 { + return current, "" + } + return current + incoming, incoming +} + +// NewClient 创建一个新的OpenAI客户端。 +func NewClient(cfg *config.OpenAIConfig, httpClient *http.Client, logger *zap.Logger) *Client { + if httpClient == nil { + httpClient = http.DefaultClient + } + if logger == nil { + logger = zap.NewNop() + } + return &Client{ + httpClient: httpClient, + config: cfg, + logger: logger, + } +} + +// UpdateConfig 动态更新OpenAI配置。 +func (c *Client) UpdateConfig(cfg *config.OpenAIConfig) { + c.config = cfg +} + +// ChatCompletion 调用 /chat/completions 接口。 +func (c *Client) ChatCompletion(ctx context.Context, payload interface{}, out interface{}) error { + if c == nil { + return fmt.Errorf("openai client is not initialized") + } + if c.config == nil { + return fmt.Errorf("openai config is nil") + } + if strings.TrimSpace(c.config.APIKey) == "" { + return fmt.Errorf("openai api key is empty") + } + if c.isClaude() { + return c.claudeChatCompletion(ctx, payload, out) + } + + baseURL := strings.TrimSuffix(c.config.BaseURL, "/") + if baseURL == "" { + baseURL = "https://api.openai.com/v1" + } + + body, err := json.Marshal(payload) + if err != nil { + return fmt.Errorf("marshal openai payload: %w", err) + } + + c.logger.Debug("sending OpenAI chat completion request", + zap.Int("payloadSizeKB", len(body)/1024)) + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, baseURL+"/chat/completions", bytes.NewReader(body)) + if err != nil { + return fmt.Errorf("build openai request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+c.config.APIKey) + + requestStart := time.Now() + resp, err := c.httpClient.Do(req) + if err != nil { + return fmt.Errorf("call openai api: %w", err) + } + defer resp.Body.Close() + + bodyChan := make(chan []byte, 1) + errChan := make(chan error, 1) + go func() { + responseBody, err := io.ReadAll(resp.Body) + if err != nil { + errChan <- err + return + } + bodyChan <- responseBody + }() + + var respBody []byte + select { + case respBody = <-bodyChan: + case err := <-errChan: + return fmt.Errorf("read openai response: %w", err) + case <-ctx.Done(): + return fmt.Errorf("read openai response timeout: %w", ctx.Err()) + case <-time.After(25 * time.Minute): + return fmt.Errorf("read openai response timeout (25m)") + } + + c.logger.Debug("received OpenAI response", + zap.Int("status", resp.StatusCode), + zap.Duration("duration", time.Since(requestStart)), + zap.Int("responseSizeKB", len(respBody)/1024), + ) + + if resp.StatusCode != http.StatusOK { + c.logger.Warn("OpenAI chat completion returned non-200", + zap.Int("status", resp.StatusCode), + zap.String("body", string(respBody)), + ) + return &APIError{ + StatusCode: resp.StatusCode, + Body: string(respBody), + } + } + + if out != nil { + if err := json.Unmarshal(respBody, out); err != nil { + c.logger.Error("failed to unmarshal OpenAI response", + zap.Error(err), + zap.String("body", string(respBody)), + ) + return fmt.Errorf("unmarshal openai response: %w", err) + } + } + + return nil +} + +// ChatCompletionStream 调用 /chat/completions 的流式模式(stream=true),并在每个 delta 到达时回调 onDelta。 +// 返回最终拼接的 content(只拼 content delta;工具调用 delta 未做处理)。 +func (c *Client) ChatCompletionStream(ctx context.Context, payload interface{}, onDelta func(delta string) error) (string, error) { + if c == nil { + return "", fmt.Errorf("openai client is not initialized") + } + if c.config == nil { + return "", fmt.Errorf("openai config is nil") + } + if strings.TrimSpace(c.config.APIKey) == "" { + return "", fmt.Errorf("openai api key is empty") + } + if c.isClaude() { + return c.claudeChatCompletionStream(ctx, payload, onDelta) + } + + baseURL := strings.TrimSuffix(c.config.BaseURL, "/") + if baseURL == "" { + baseURL = "https://api.openai.com/v1" + } + + body, err := json.Marshal(payload) + if err != nil { + return "", fmt.Errorf("marshal openai payload: %w", err) + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, baseURL+"/chat/completions", bytes.NewReader(body)) + if err != nil { + return "", fmt.Errorf("build openai request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+c.config.APIKey) + + requestStart := time.Now() + resp, err := c.httpClient.Do(req) + if err != nil { + return "", fmt.Errorf("call openai api: %w", err) + } + defer resp.Body.Close() + + // 非200:读完 body 返回 + if resp.StatusCode != http.StatusOK { + respBody, readErr := io.ReadAll(resp.Body) + if readErr != nil { + c.logger.Warn("failed to read OpenAI error response body", zap.Error(readErr)) + } + return "", &APIError{ + StatusCode: resp.StatusCode, + Body: string(respBody), + } + } + + type streamDelta struct { + // OpenAI 兼容流式通常使用 content;但部分兼容实现可能用 text。 + Content string `json:"content,omitempty"` + Text string `json:"text,omitempty"` + } + type streamChoice struct { + Delta streamDelta `json:"delta"` + FinishReason *string `json:"finish_reason,omitempty"` + } + type streamResponse struct { + ID string `json:"id,omitempty"` + Choices []streamChoice `json:"choices"` + Error *struct { + Message string `json:"message"` + Type string `json:"type"` + } `json:"error,omitempty"` + } + + reader := bufio.NewReader(resp.Body) + var full strings.Builder + fullText := "" + + // 典型 SSE 结构: + // data: {...}\n\n + // data: [DONE]\n\n + for { + line, readErr := reader.ReadString('\n') + if readErr != nil { + if readErr == io.EOF { + break + } + return full.String(), fmt.Errorf("read openai stream: %w", readErr) + } + trimmed := strings.TrimSpace(line) + if trimmed == "" { + continue + } + if !strings.HasPrefix(trimmed, "data:") { + continue + } + dataStr := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) + if dataStr == "[DONE]" { + break + } + + var chunk streamResponse + if err := json.Unmarshal([]byte(dataStr), &chunk); err != nil { + // 解析失败跳过(兼容各种兼容层的差异) + continue + } + if chunk.Error != nil && strings.TrimSpace(chunk.Error.Message) != "" { + return full.String(), fmt.Errorf("openai stream error: %s", chunk.Error.Message) + } + if len(chunk.Choices) == 0 { + continue + } + + delta := chunk.Choices[0].Delta.Content + if delta == "" { + delta = chunk.Choices[0].Delta.Text + } + if delta == "" { + continue + } + + var deltaOut string + fullText, deltaOut = normalizeStreamingDelta(fullText, delta) + if deltaOut == "" { + continue + } + full.WriteString(deltaOut) + if onDelta != nil { + if err := onDelta(deltaOut); err != nil { + return full.String(), err + } + } + } + + c.logger.Debug("received OpenAI stream completion", + zap.Duration("duration", time.Since(requestStart)), + zap.Int("contentLen", full.Len()), + ) + + return full.String(), nil +} + +// StreamToolCall 流式工具调用的累积结果(arguments 以字符串形式拼接,留给上层再解析为 JSON)。 +type StreamToolCall struct { + Index int + ID string + Type string + FunctionName string + FunctionArgsStr string +} + +// ChatCompletionStreamWithToolCalls 流式模式:同时把 content delta 实时回调,并在结束后返回 tool_calls 和 finish_reason。 +func (c *Client) ChatCompletionStreamWithToolCalls( + ctx context.Context, + payload interface{}, + onContentDelta func(delta string) error, +) (string, []StreamToolCall, string, error) { + if c == nil { + return "", nil, "", fmt.Errorf("openai client is not initialized") + } + if c.config == nil { + return "", nil, "", fmt.Errorf("openai config is nil") + } + if strings.TrimSpace(c.config.APIKey) == "" { + return "", nil, "", fmt.Errorf("openai api key is empty") + } + if c.isClaude() { + return c.claudeChatCompletionStreamWithToolCalls(ctx, payload, onContentDelta) + } + + baseURL := strings.TrimSuffix(c.config.BaseURL, "/") + if baseURL == "" { + baseURL = "https://api.openai.com/v1" + } + + body, err := json.Marshal(payload) + if err != nil { + return "", nil, "", fmt.Errorf("marshal openai payload: %w", err) + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, baseURL+"/chat/completions", bytes.NewReader(body)) + if err != nil { + return "", nil, "", fmt.Errorf("build openai request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+c.config.APIKey) + + requestStart := time.Now() + resp, err := c.httpClient.Do(req) + if err != nil { + return "", nil, "", fmt.Errorf("call openai api: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + respBody, readErr := io.ReadAll(resp.Body) + if readErr != nil { + c.logger.Warn("failed to read OpenAI error response body", zap.Error(readErr)) + } + return "", nil, "", &APIError{ + StatusCode: resp.StatusCode, + Body: string(respBody), + } + } + + // delta tool_calls 的增量结构 + type toolCallFunctionDelta struct { + Name string `json:"name,omitempty"` + Arguments string `json:"arguments,omitempty"` + } + type toolCallDelta struct { + Index int `json:"index,omitempty"` + ID string `json:"id,omitempty"` + Type string `json:"type,omitempty"` + Function toolCallFunctionDelta `json:"function,omitempty"` + } + type streamDelta2 struct { + Content string `json:"content,omitempty"` + Text string `json:"text,omitempty"` + ToolCalls []toolCallDelta `json:"tool_calls,omitempty"` + } + type streamChoice2 struct { + Delta streamDelta2 `json:"delta"` + FinishReason *string `json:"finish_reason,omitempty"` + } + type streamResponse2 struct { + Choices []streamChoice2 `json:"choices"` + Error *struct { + Message string `json:"message"` + Type string `json:"type"` + } `json:"error,omitempty"` + } + + type toolCallAccum struct { + id string + typ string + name string + args strings.Builder + } + toolCallAccums := make(map[int]*toolCallAccum) + + reader := bufio.NewReader(resp.Body) + var full strings.Builder + fullText := "" + finishReason := "" + + for { + line, readErr := reader.ReadString('\n') + if readErr != nil { + if readErr == io.EOF { + break + } + return full.String(), nil, finishReason, fmt.Errorf("read openai stream: %w", readErr) + } + trimmed := strings.TrimSpace(line) + if trimmed == "" { + continue + } + if !strings.HasPrefix(trimmed, "data:") { + continue + } + dataStr := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) + if dataStr == "[DONE]" { + break + } + + var chunk streamResponse2 + if err := json.Unmarshal([]byte(dataStr), &chunk); err != nil { + // 兼容:解析失败跳过 + continue + } + if chunk.Error != nil && strings.TrimSpace(chunk.Error.Message) != "" { + return full.String(), nil, finishReason, fmt.Errorf("openai stream error: %s", chunk.Error.Message) + } + if len(chunk.Choices) == 0 { + continue + } + + choice := chunk.Choices[0] + if choice.FinishReason != nil && strings.TrimSpace(*choice.FinishReason) != "" { + finishReason = strings.TrimSpace(*choice.FinishReason) + } + + delta := choice.Delta + + content := delta.Content + if content == "" { + content = delta.Text + } + if content != "" { + var contentOut string + fullText, contentOut = normalizeStreamingDelta(fullText, content) + if contentOut != "" { + full.WriteString(contentOut) + if onContentDelta != nil { + if err := onContentDelta(contentOut); err != nil { + return full.String(), nil, finishReason, err + } + } + } + } + + if len(delta.ToolCalls) > 0 { + for _, tc := range delta.ToolCalls { + acc, ok := toolCallAccums[tc.Index] + if !ok { + acc = &toolCallAccum{} + toolCallAccums[tc.Index] = acc + } + if tc.ID != "" { + acc.id = tc.ID + } + if tc.Type != "" { + acc.typ = tc.Type + } + if tc.Function.Name != "" { + acc.name = tc.Function.Name + } + if tc.Function.Arguments != "" { + acc.args.WriteString(tc.Function.Arguments) + } + } + } + } + + // 组装 tool calls + indices := make([]int, 0, len(toolCallAccums)) + for idx := range toolCallAccums { + indices = append(indices, idx) + } + // 手写简单排序(避免额外 import) + for i := 0; i < len(indices); i++ { + for j := i + 1; j < len(indices); j++ { + if indices[j] < indices[i] { + indices[i], indices[j] = indices[j], indices[i] + } + } + } + + toolCalls := make([]StreamToolCall, 0, len(indices)) + for _, idx := range indices { + acc := toolCallAccums[idx] + tc := StreamToolCall{ + Index: idx, + ID: acc.id, + Type: acc.typ, + FunctionName: acc.name, + FunctionArgsStr: acc.args.String(), + } + toolCalls = append(toolCalls, tc) + } + + c.logger.Debug("received OpenAI stream completion (tool_calls)", + zap.Duration("duration", time.Since(requestStart)), + zap.Int("contentLen", full.Len()), + zap.Int("toolCalls", len(toolCalls)), + zap.String("finishReason", finishReason), + ) + + if strings.TrimSpace(finishReason) == "" { + finishReason = "stop" + } + + return full.String(), toolCalls, finishReason, nil +} + +// ModelsListResponse 表示 OpenAI 兼容 GET /models 响应。 +type ModelsListResponse struct { + Object string `json:"object"` + Data []struct { + ID string `json:"id"` + Object string `json:"object,omitempty"` + OwnedBy string `json:"owned_by,omitempty"` + } `json:"data"` +} + +// ListModels 调用 GET {baseURL}/models 获取可用模型 id 列表(按字典序)。 +func (c *Client) ListModels(ctx context.Context) ([]string, error) { + if c == nil { + return nil, fmt.Errorf("openai client is not initialized") + } + if c.config == nil { + return nil, fmt.Errorf("openai config is nil") + } + if strings.TrimSpace(c.config.APIKey) == "" { + return nil, fmt.Errorf("openai api key is empty") + } + if c.isClaude() { + return nil, fmt.Errorf("claude provider does not support models list API") + } + + baseURL := strings.TrimSuffix(c.config.BaseURL, "/") + if baseURL == "" { + baseURL = "https://api.openai.com/v1" + } + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL+"/models", nil) + if err != nil { + return nil, fmt.Errorf("build openai models request: %w", err) + } + req.Header.Set("Authorization", "Bearer "+c.config.APIKey) + + resp, err := c.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("call openai models api: %w", err) + } + defer resp.Body.Close() + + respBody, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("read openai models response: %w", err) + } + if resp.StatusCode != http.StatusOK { + return nil, &APIError{ + StatusCode: resp.StatusCode, + Body: string(respBody), + } + } + + var list ModelsListResponse + if err := json.Unmarshal(respBody, &list); err != nil { + return nil, fmt.Errorf("decode openai models response: %w", err) + } + + seen := make(map[string]struct{}, len(list.Data)) + models := make([]string, 0, len(list.Data)) + for _, item := range list.Data { + id := strings.TrimSpace(item.ID) + if id == "" { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + models = append(models, id) + } + sort.Strings(models) + if len(models) == 0 { + return nil, fmt.Errorf("models list is empty") + } + return models, nil +} diff --git a/internal/openai/reasoning_payload.go b/internal/openai/reasoning_payload.go new file mode 100644 index 00000000..2110c90d --- /dev/null +++ b/internal/openai/reasoning_payload.go @@ -0,0 +1,117 @@ +package openai + +import ( + "strings" + + "github.com/bytedance/sonic" +) + +// reasoningPayloadKeys are OpenAI-compatible root fields that enable "thinking" / +// extended-reasoning modes on gateways such as DashScope/Qwen and MiniMax. +var reasoningPayloadKeys = []string{ + "thinking", + "reasoning_effort", + "output_config", + "reasoning", +} + +// StripReasoningFromChatCompletionBody removes thinking / reasoning fields from a +// chat-completions JSON body. +func StripReasoningFromChatCompletionBody(rawBody []byte) ([]byte, error) { + var payload map[string]any + if err := sonic.Unmarshal(rawBody, &payload); err != nil { + return rawBody, nil + } + if !stripReasoningFields(payload) { + return rawBody, nil + } + out, err := sonic.Marshal(payload) + if err != nil { + return rawBody, err + } + return out, nil +} + +// StripReasoningIfForcedToolChoice removes thinking / reasoning fields when the +// request sets tool_choice to "required" or an object. Several providers reject +// that combination (e.g. DashScope: "tool_choice does not support being set to +// required or object in thinking mode"). +func StripReasoningIfForcedToolChoice(rawBody []byte) ([]byte, error) { + var payload map[string]any + if err := sonic.Unmarshal(rawBody, &payload); err != nil { + return rawBody, nil + } + if !forcedToolChoiceIncompatibleWithThinking(payload) { + return rawBody, nil + } + if !stripReasoningFields(payload) { + return rawBody, nil + } + out, err := sonic.Marshal(payload) + if err != nil { + return rawBody, err + } + return out, nil +} + +// StripToolChoiceForThinkingMode removes tool_choice while preserving tools and +// thinking fields. DeepSeek thinking mode can use tools, but rejects the +// tool_choice parameter itself on some agent requests. +func StripToolChoiceForThinkingMode(rawBody []byte) ([]byte, error) { + var payload map[string]any + if err := sonic.Unmarshal(rawBody, &payload); err != nil { + return rawBody, nil + } + if !thinkingModeEnabledByPayload(payload) { + return rawBody, nil + } + if _, ok := payload["tool_choice"]; !ok { + return rawBody, nil + } + delete(payload, "tool_choice") + out, err := sonic.Marshal(payload) + if err != nil { + return rawBody, err + } + return out, nil +} + +func stripReasoningFields(payload map[string]any) bool { + changed := false + for _, key := range reasoningPayloadKeys { + if _, ok := payload[key]; ok { + delete(payload, key) + changed = true + } + } + return changed +} + +func forcedToolChoiceIncompatibleWithThinking(payload map[string]any) bool { + tc, ok := payload["tool_choice"] + if !ok || tc == nil { + return false + } + switch v := tc.(type) { + case string: + return v == "required" + case map[string]any: + return true + default: + return false + } +} + +func thinkingModeEnabledByPayload(payload map[string]any) bool { + thinking, ok := payload["thinking"] + if !ok || thinking == nil { + // DeepSeek enables thinking by default unless explicitly disabled. + return true + } + if m, ok := thinking.(map[string]any); ok { + if typ, ok := m["type"].(string); ok && strings.EqualFold(strings.TrimSpace(typ), "disabled") { + return false + } + } + return true +} diff --git a/internal/openai/reasoning_payload_test.go b/internal/openai/reasoning_payload_test.go new file mode 100644 index 00000000..9f7dccfe --- /dev/null +++ b/internal/openai/reasoning_payload_test.go @@ -0,0 +1,250 @@ +package openai + +import ( + "io" + "net/http" + "strings" + "testing" + + "cyberstrike-ai/internal/config" +) + +func TestStripReasoningFromChatCompletionBody(t *testing.T) { + in := []byte(`{"model":"deepseek-chat","messages":[],"thinking":{"type":"enabled"},"reasoning_effort":"high"}`) + out, err := StripReasoningFromChatCompletionBody(in) + if err != nil { + t.Fatal(err) + } + s := string(out) + if strings.Contains(s, "thinking") || strings.Contains(s, "reasoning_effort") { + t.Fatalf("expected reasoning fields stripped, got %s", s) + } + if !strings.Contains(s, `"model":"deepseek-chat"`) { + t.Fatalf("expected model preserved, got %s", s) + } + + plain := []byte(`{"model":"gpt-4o","messages":[]}`) + out2, err := StripReasoningFromChatCompletionBody(plain) + if err != nil { + t.Fatal(err) + } + if string(out2) != string(plain) { + t.Fatalf("expected unchanged payload, got %s", out2) + } +} + +func TestStripReasoningIfForcedToolChoice(t *testing.T) { + cases := []struct { + name string + in string + strip bool + contain string + }{ + { + name: "required strips thinking", + in: `{"model":"minimax","messages":[],"thinking":{"type":"enabled"},"tool_choice":"required","tools":[]}`, + strip: true, + }, + { + name: "object tool_choice strips thinking", + in: `{"model":"qwen","messages":[],"thinking":{"type":"enabled"},"tool_choice":{"type":"function","function":{"name":"respond"}}}`, + strip: true, + }, + { + name: "auto keeps thinking", + in: `{"model":"qwen","messages":[],"thinking":{"type":"enabled"},"tool_choice":"auto"}`, + strip: false, + contain: "thinking", + }, + { + name: "no tool_choice keeps thinking", + in: `{"model":"qwen","messages":[],"thinking":{"type":"enabled"}}`, + strip: false, + contain: "thinking", + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + out, err := StripReasoningIfForcedToolChoice([]byte(tc.in)) + if err != nil { + t.Fatal(err) + } + s := string(out) + hasThinking := strings.Contains(s, "thinking") + if tc.strip && hasThinking { + t.Fatalf("expected thinking stripped, got %s", s) + } + if !tc.strip && tc.contain != "" && !strings.Contains(s, tc.contain) { + t.Fatalf("expected %q in %s", tc.contain, s) + } + if !tc.strip && string(out) != tc.in { + t.Fatalf("expected unchanged payload, got %s", s) + } + }) + } +} + +func TestStripToolChoiceForThinkingMode(t *testing.T) { + cases := []struct { + name string + in string + wantToolChoice bool + wantThinking bool + }{ + { + name: "enabled thinking removes tool_choice", + in: `{"model":"deepseek-v4","messages":[],"thinking":{"type":"enabled"},"tool_choice":"required","tools":[{"type":"function","function":{"name":"scan"}}]}`, + wantToolChoice: false, + wantThinking: true, + }, + { + name: "default thinking removes tool_choice", + in: `{"model":"deepseek-v4","messages":[],"tool_choice":"auto","tools":[]}`, + wantToolChoice: false, + wantThinking: false, + }, + { + name: "disabled thinking keeps tool_choice", + in: `{"model":"deepseek-v4","messages":[],"thinking":{"type":"disabled"},"tool_choice":"required","tools":[]}`, + wantToolChoice: true, + wantThinking: true, + }, + { + name: "no tool_choice unchanged", + in: `{"model":"deepseek-v4","messages":[],"thinking":{"type":"enabled"},"tools":[]}`, + wantToolChoice: false, + wantThinking: true, + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + out, err := StripToolChoiceForThinkingMode([]byte(tc.in)) + if err != nil { + t.Fatal(err) + } + s := string(out) + if strings.Contains(s, "tool_choice") != tc.wantToolChoice { + t.Fatalf("tool_choice presence mismatch, got %s", s) + } + if strings.Contains(s, "thinking") != tc.wantThinking { + t.Fatalf("thinking presence mismatch, got %s", s) + } + if !strings.Contains(s, "tools") { + t.Fatalf("expected tools preserved, got %s", s) + } + }) + } +} + +func TestReasoningToolChoiceCompatRoundTripper(t *testing.T) { + var gotBody string + rt := &reasoningToolChoiceCompatRoundTripper{ + base: roundTripperFunc(func(req *http.Request) (*http.Response, error) { + b, _ := io.ReadAll(req.Body) + gotBody = string(b) + return &http.Response{ + StatusCode: 200, + Body: io.NopCloser(strings.NewReader(`{"choices":[{"message":{"content":"ok"}}]}`)), + Header: http.Header{"Content-Type": []string{"application/json"}}, + }, nil + }), + } + req, err := http.NewRequest(http.MethodPost, "https://example.com/v1/chat/completions", strings.NewReader( + `{"model":"m","thinking":{"type":"enabled"},"tool_choice":"required","messages":[]}`, + )) + if err != nil { + t.Fatal(err) + } + _, err = rt.RoundTrip(req) + if err != nil { + t.Fatal(err) + } + if strings.Contains(gotBody, "thinking") { + t.Fatalf("expected thinking stripped in transit, got %s", gotBody) + } + if !strings.Contains(gotBody, `"tool_choice":"required"`) { + t.Fatalf("expected tool_choice preserved, got %s", gotBody) + } +} + +func TestReasoningToolChoiceCompatRoundTripperDeepSeek(t *testing.T) { + var gotBody string + rt := &reasoningToolChoiceCompatRoundTripper{ + cfg: &config.OpenAIConfig{ + BaseURL: "https://api.deepseek.com/v1", + Model: "deepseek-v4", + }, + base: roundTripperFunc(func(req *http.Request) (*http.Response, error) { + b, _ := io.ReadAll(req.Body) + gotBody = string(b) + return &http.Response{ + StatusCode: 200, + Body: io.NopCloser(strings.NewReader(`{"choices":[{"message":{"content":"ok"}}]}`)), + Header: http.Header{"Content-Type": []string{"application/json"}}, + }, nil + }), + } + req, err := http.NewRequest(http.MethodPost, "https://api.deepseek.com/v1/chat/completions", strings.NewReader( + `{"model":"deepseek-v4","thinking":{"type":"enabled"},"tool_choice":"required","tools":[],"messages":[]}`, + )) + if err != nil { + t.Fatal(err) + } + _, err = rt.RoundTrip(req) + if err != nil { + t.Fatal(err) + } + if strings.Contains(gotBody, "tool_choice") { + t.Fatalf("expected DeepSeek tool_choice stripped in transit, got %s", gotBody) + } + if !strings.Contains(gotBody, "thinking") { + t.Fatalf("expected thinking preserved for DeepSeek, got %s", gotBody) + } + if !strings.Contains(gotBody, "tools") { + t.Fatalf("expected tools preserved for DeepSeek, got %s", gotBody) + } +} + +func TestReasoningToolChoiceCompatRoundTripperDeepSeekEndpointWinsOverProfile(t *testing.T) { + var gotBody string + rt := &reasoningToolChoiceCompatRoundTripper{ + cfg: &config.OpenAIConfig{ + BaseURL: "https://api.deepseek.com/v1", + Model: "deepseek-v4-flash", + Reasoning: config.OpenAIReasoningConfig{ + Profile: "openai_compat", + }, + }, + base: roundTripperFunc(func(req *http.Request) (*http.Response, error) { + b, _ := io.ReadAll(req.Body) + gotBody = string(b) + return &http.Response{ + StatusCode: 200, + Body: io.NopCloser(strings.NewReader(`{"choices":[{"message":{"content":"ok"}}]}`)), + Header: http.Header{"Content-Type": []string{"application/json"}}, + }, nil + }), + } + req, err := http.NewRequest(http.MethodPost, "https://api.deepseek.com/v1/chat/completions", strings.NewReader( + `{"model":"deepseek-v4-flash","tool_choice":"required","tools":[],"messages":[]}`, + )) + if err != nil { + t.Fatal(err) + } + _, err = rt.RoundTrip(req) + if err != nil { + t.Fatal(err) + } + if strings.Contains(gotBody, "tool_choice") { + t.Fatalf("expected DeepSeek tool_choice stripped despite openai_compat profile, got %s", gotBody) + } + if !strings.Contains(gotBody, "tools") { + t.Fatalf("expected tools preserved for DeepSeek, got %s", gotBody) + } +} + +type roundTripperFunc func(*http.Request) (*http.Response, error) + +func (f roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} diff --git a/internal/openai/reasoning_tool_choice_compat.go b/internal/openai/reasoning_tool_choice_compat.go new file mode 100644 index 00000000..fd703222 --- /dev/null +++ b/internal/openai/reasoning_tool_choice_compat.go @@ -0,0 +1,69 @@ +package openai + +import ( + "bytes" + "io" + "net/http" + "strconv" + "strings" + + "cyberstrike-ai/internal/config" +) + +// reasoningToolChoiceCompatRoundTripper strips thinking/reasoning fields from +// chat/completions requests that force tool_choice, which some gateways reject +// when thinking mode is enabled on the same request. +type reasoningToolChoiceCompatRoundTripper struct { + base http.RoundTripper + cfg *config.OpenAIConfig +} + +func (rt *reasoningToolChoiceCompatRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + if rt == nil || rt.base == nil || req == nil || req.Body == nil { + if rt != nil && rt.base != nil { + return rt.base.RoundTrip(req) + } + return http.DefaultTransport.RoundTrip(req) + } + if req.Method != http.MethodPost || !strings.HasSuffix(req.URL.Path, "/chat/completions") { + return rt.base.RoundTrip(req) + } + + body, err := io.ReadAll(req.Body) + _ = req.Body.Close() + if err != nil { + return nil, err + } + + patched := body + var perr error + if isDeepSeekToolChoiceCompatProfile(rt.cfg) { + patched, perr = StripToolChoiceForThinkingMode(body) + } else { + patched, perr = StripReasoningIfForcedToolChoice(body) + } + if perr != nil { + patched = body + } + req.Body = io.NopCloser(bytes.NewReader(patched)) + req.ContentLength = int64(len(patched)) + req.Header.Set("Content-Length", strconv.Itoa(len(patched))) + return rt.base.RoundTrip(req) +} + +func isDeepSeekToolChoiceCompatProfile(cfg *config.OpenAIConfig) bool { + if cfg == nil { + return false + } + if cfg.IsDeepSeekEndpointOrModel() { + return true + } + profile := strings.ToLower(strings.TrimSpace(cfg.Reasoning.ProfileEffective())) + if profile == "deepseek" || profile == "deepseek_compat" { + return true + } + if profile != "" && profile != "auto" { + return false + } + return false +} diff --git a/internal/openai/sse_stream.go b/internal/openai/sse_stream.go new file mode 100644 index 00000000..a86d6306 --- /dev/null +++ b/internal/openai/sse_stream.go @@ -0,0 +1,20 @@ +package openai + +// SSEAccumulatedKey 为 SSE progress 事件 data 中的服务端权威流式全文快照字段。 +// 前端应优先用该字段更新 buffer,避免对 delta 二次 normalize 导致叠字。 +const SSEAccumulatedKey = "accumulated" + +// WithSSEAccumulated 在 progress data 中附带当前流式累计全文(权威快照)。 +func WithSSEAccumulated(data map[string]interface{}, accumulated string) map[string]interface{} { + if data == nil { + data = make(map[string]interface{}, 1) + } + data[SSEAccumulatedKey] = accumulated + return data +} + +// NormalizeStreamingDelta 将可能是“累计片段/重发片段”的内容归一化为“纯增量”。 +// 与 unexported normalizeStreamingDelta 相同,供 agent / multiagent 等包在发 SSE 前累计正文。 +func NormalizeStreamingDelta(current, incoming string) (next, delta string) { + return normalizeStreamingDelta(current, incoming) +} diff --git a/internal/openai/summarization_diag.go b/internal/openai/summarization_diag.go new file mode 100644 index 00000000..44465145 --- /dev/null +++ b/internal/openai/summarization_diag.go @@ -0,0 +1,108 @@ +package openai + +import ( + "bytes" + "io" + "net/http" + "strings" + + "github.com/bytedance/sonic" + "go.uber.org/zap" +) + +// SummarizationRequestHeader marks chat/completion requests issued by Eino summarization +// middleware (via model.WithExtraHeader). The diagnostic transport logs empty-choices bodies +// only for these requests so main-agent traffic stays quiet. +const SummarizationRequestHeader = "X-CyberStrike-Summarization" + +const summarizationDiagBodyMaxBytes = 8192 + +// AttachSummarizationDiagTransport wraps client.Transport to log raw API bodies when +// summarization receives HTTP 200 with an empty choices array. +func AttachSummarizationDiagTransport(client *http.Client, logger *zap.Logger) { + if client == nil || logger == nil { + return + } + base := client.Transport + if base == nil { + base = http.DefaultTransport + } + client.Transport = &summarizationDiagRoundTripper{base: base, logger: logger} +} + +type summarizationDiagRoundTripper struct { + base http.RoundTripper + logger *zap.Logger +} + +func (rt *summarizationDiagRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + resp, err := rt.base.RoundTrip(req) + if err != nil || resp == nil || resp.Body == nil { + return resp, err + } + if !isSummarizationRequest(req) || !strings.Contains(strings.ToLower(resp.Header.Get("Content-Type")), "json") { + return resp, err + } + + body, readErr := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if readErr != nil { + resp.Body = io.NopCloser(bytes.NewReader(nil)) + return resp, err + } + resp.Body = io.NopCloser(bytes.NewReader(body)) + resp.ContentLength = int64(len(body)) + + if rt.logger != nil && resp.StatusCode >= http.StatusBadRequest { + rt.logger.Warn("eino summarization: API request rejected", + zap.Int("status", resp.StatusCode), + zap.String("request_id", responseRequestID(resp)), + zap.Int("response_bytes", len(body)), + zap.String("raw_body", truncateForLog(string(body), summarizationDiagBodyMaxBytes)), + ) + } else if rt.logger != nil && summarizationResponseEmptyChoices(body) { + rt.logger.Warn("eino summarization: API returned empty choices", + zap.Int("status", resp.StatusCode), + zap.String("request_id", responseRequestID(resp)), + zap.Int("response_bytes", len(body)), + zap.String("raw_body", truncateForLog(string(body), summarizationDiagBodyMaxBytes)), + ) + } + return resp, err +} + +func responseRequestID(resp *http.Response) string { + if resp == nil { + return "" + } + for _, key := range []string{"x-request-id", "request-id", "x-trace-id"} { + if value := strings.TrimSpace(resp.Header.Get(key)); value != "" { + return value + } + } + return "" +} + +func isSummarizationRequest(req *http.Request) bool { + if req == nil { + return false + } + return strings.TrimSpace(req.Header.Get(SummarizationRequestHeader)) == "1" +} + +func summarizationResponseEmptyChoices(body []byte) bool { + var parsed struct { + Choices []any `json:"choices"` + } + if err := sonic.Unmarshal(body, &parsed); err != nil { + return false + } + return len(parsed.Choices) == 0 +} + +func truncateForLog(s string, maxBytes int) string { + if maxBytes <= 0 || len(s) <= maxBytes { + return s + } + return s[:maxBytes] + "…(truncated)" +} diff --git a/internal/openai/summarization_diag_test.go b/internal/openai/summarization_diag_test.go new file mode 100644 index 00000000..753a61ae --- /dev/null +++ b/internal/openai/summarization_diag_test.go @@ -0,0 +1,47 @@ +package openai + +import ( + "io" + "net/http" + "strings" + "testing" + + "go.uber.org/zap" +) + +type staticRoundTripper struct { + status int + body string +} + +func (s *staticRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: s.status, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(s.body)), + }, nil +} + +func TestSummarizationResponseEmptyChoices(t *testing.T) { + if !summarizationResponseEmptyChoices([]byte(`{"choices":[]}`)) { + t.Fatal("expected empty choices") + } + if summarizationResponseEmptyChoices([]byte(`{"choices":[{"index":0}]}`)) { + t.Fatal("expected non-empty choices") + } +} + +func TestSummarizationDiagRoundTripper_SkipsWithoutHeader(t *testing.T) { + client := &http.Client{ + Transport: &summarizationDiagRoundTripper{ + base: &staticRoundTripper{status: 200, body: `{"choices":[]}`}, + logger: zap.NewNop(), + }, + } + req, _ := http.NewRequest(http.MethodPost, "https://example.com/v1/chat/completions", nil) + resp, err := client.Do(req) + if err != nil { + t.Fatal(err) + } + _ = resp.Body.Close() +} diff --git a/internal/tooloutput/spill.go b/internal/tooloutput/spill.go new file mode 100644 index 00000000..1352e219 --- /dev/null +++ b/internal/tooloutput/spill.go @@ -0,0 +1,292 @@ +// Package tooloutput spills oversized tool stdout/results to local files under +// the reduction cache tree (tmp/reduction/...), so agents can read_file the +// full text after context truncation. +package tooloutput + +import ( + "fmt" + "os" + "path/filepath" + "strings" + "sync" + "unicode/utf8" + + "github.com/google/uuid" +) + +const ( + defaultRootDir = "tmp/reduction" + readFileHint = "read_file" +) + +// SpillOpts scopes where a trunc file is written (mirrors reduction RootDir layout). +type SpillOpts struct { + RootDir string // reduction_root_dir or empty → tmp/reduction + ProjectID string + ConversationID string + ExecutionID string // preferred file name; empty → uuid +} + +// SessionRoot returns the conversation/project-scoped reduction cache root. +func SessionRoot(configuredBase, projectID, conversationID string) string { + base := strings.TrimSpace(configuredBase) + if base == "" { + base = defaultRootDir + } + if pid := strings.TrimSpace(projectID); pid != "" { + return filepath.Join(base, "projects", sanitizeSegment(pid)) + } + conv := strings.TrimSpace(conversationID) + if conv == "" { + conv = "default" + } + return filepath.Join(base, "conversations", sanitizeSegment(conv)) +} + +// WriteTruncFile writes full content under {sessionRoot}/trunc/{id} and returns +// an absolute path suitable for read_file. +func WriteTruncFile(opts SpillOpts, content string) (string, error) { + session := SessionRoot(opts.RootDir, opts.ProjectID, opts.ConversationID) + id := strings.TrimSpace(opts.ExecutionID) + if id == "" { + id = uuid.NewString() + } + id = sanitizeSegment(id) + dir := filepath.Join(session, "trunc") + if err := os.MkdirAll(dir, 0o755); err != nil { + return "", fmt.Errorf("mkdir tool output trunc dir: %w", err) + } + path := filepath.Join(dir, id) + if abs, err := filepath.Abs(path); err == nil { + path = abs + } + if err := os.WriteFile(path, []byte(content), 0o600); err != nil { + return "", fmt.Errorf("write tool output trunc file: %w", err) + } + return path, nil +} + +// BoundWithSpill truncates full text into a notice after +// spilling the original to disk. The returned string is always ≤ maxBytes when +// maxBytes > 0. On spill failure it falls back to a prefix + marker (no path). +func BoundWithSpill(full string, maxBytes int, opts SpillOpts) string { + if maxBytes <= 0 || len(full) <= maxBytes { + return full + } + path, err := WriteTruncFile(opts, full) + if err != nil { + return boundPrefixOnly(full, maxBytes, len(full), "") + } + return FormatPersistedOutput(full, path, maxBytes) +} + +// FormatPersistedOutput builds a reduction-compatible notice with head/tail +// previews that fits in maxBytes. +func FormatPersistedOutput(full, filePath string, maxBytes int) string { + return formatPersisted(len(full), filePath, full, maxBytes) +} + +// FormatPersistedFromFile builds the notice using previews read from an already +// spilled file (streaming collectors that never kept the full string in memory). +func FormatPersistedFromFile(filePath string, originalSize, maxBytes int) string { + previewSrc := "" + if data, err := os.ReadFile(filePath); err == nil { + previewSrc = string(data) + if originalSize <= 0 { + originalSize = len(data) + } + } + return formatPersisted(originalSize, filePath, previewSrc, maxBytes) +} + +func formatPersisted(originalSize int, filePath, previewSrc string, maxBytes int) string { + if maxBytes <= 0 { + maxBytes = 12000 + } + // Always keep the absolute path readable for read_file, even under tight budgets. + minimal := fmt.Sprintf( + "\nOutput too large (%d). Full output saved to: %s\nUse %s to read.\n", + originalSize, filePath, readFileHint, + ) + if len(minimal) > maxBytes { + core := fmt.Sprintf("Full output saved to: %s", filePath) + if len(core) <= maxBytes { + return core + } + // Path longer than budget: keep as much of the path as possible after a short prefix. + prefix := "Full output saved to: " + suffix := "" + room := maxBytes - len(prefix) - len(suffix) + if room <= 0 { + return clampPrefix(core, maxBytes) + } + return prefix + clampSuffix(filePath, room) + suffix + } + + previewBudget := maxBytes - len(minimal) + 32 // approximate room beyond minimal shell + if previewBudget > 4000 { + previewBudget = 4000 + } + if previewBudget < 0 { + previewBudget = 0 + } + for previewBudget >= 0 { + half := previewBudget / 2 + head := clampPrefix(previewSrc, half) + tail := clampSuffix(previewSrc, previewBudget-half) + notice := fmt.Sprintf( + "\nOutput too large (%d). Full output saved to: %s\nUse %s with offset/limit to read parts of the file.\nPreview (first %d):\n%s\n\nPreview (last %d):\n%s\n\n", + originalSize, filePath, readFileHint, len(head), head, len(tail), tail, + ) + if len(notice) <= maxBytes { + return notice + } + if previewBudget == 0 { + return minimal + } + previewBudget = previewBudget * 3 / 4 + } + return minimal +} + +func boundPrefixOnly(full string, maxBytes, originalSize int, filePath string) string { + marker := fmt.Sprintf("\n\n...[tool output truncated: original %d bytes, kept %d bytes]...", originalSize, maxBytes) + if filePath != "" { + marker = fmt.Sprintf("\n\n...[tool output truncated: original %d bytes, kept %d bytes; full output: %s]...", originalSize, maxBytes, filePath) + } + budget := maxBytes - len(marker) + if budget < 0 { + return clampPrefix(marker, maxBytes) + } + return clampPrefix(full, budget) + marker +} + +func clampPrefix(s string, n int) string { + if n <= 0 { + return "" + } + if len(s) <= n { + return s + } + for n > 0 && !utf8.RuneStart(s[n]) { + n-- + } + return s[:n] +} + +func clampSuffix(s string, n int) string { + if n <= 0 { + return "" + } + if len(s) <= n { + return s + } + start := len(s) - n + for start < len(s) && !utf8.RuneStart(s[start]) { + start++ + } + return s[start:] +} + +func sanitizeSegment(s string) string { + s = strings.TrimSpace(s) + if s == "" { + return "default" + } + s = strings.ReplaceAll(s, string(filepath.Separator), "-") + s = strings.ReplaceAll(s, "/", "-") + s = strings.ReplaceAll(s, "\\", "-") + s = strings.ReplaceAll(s, "..", "__") + if len(s) > 180 { + s = s[:180] + } + return s +} + +// Tee writes every byte to a trunc file while callers keep only a bounded +// in-memory prefix. Safe for concurrent stdout/stderr writers. +type Tee struct { + mu sync.Mutex + opts SpillOpts + file *os.File + path string + err error + open bool +} + +// NewTee prepares a lazy spill file (created on first Write). +func NewTee(opts SpillOpts) *Tee { + return &Tee{opts: opts} +} + +// Write appends to the spill file, creating it on first use. +func (t *Tee) Write(p []byte) (int, error) { + if t == nil { + return len(p), nil + } + t.mu.Lock() + defer t.mu.Unlock() + if err := t.ensureOpenLocked(); err != nil { + return len(p), nil // best-effort: never fail the tool pipe + } + if t.file == nil { + return len(p), nil + } + _, _ = t.file.Write(p) + return len(p), nil +} + +func (t *Tee) ensureOpenLocked() error { + if t.open || t.err != nil { + return t.err + } + t.open = true + session := SessionRoot(t.opts.RootDir, t.opts.ProjectID, t.opts.ConversationID) + id := strings.TrimSpace(t.opts.ExecutionID) + if id == "" { + id = uuid.NewString() + } + id = sanitizeSegment(id) + dir := filepath.Join(session, "trunc") + if err := os.MkdirAll(dir, 0o755); err != nil { + t.err = err + return err + } + path := filepath.Join(dir, id) + if abs, err := filepath.Abs(path); err == nil { + path = abs + } + f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600) + if err != nil { + t.err = err + return err + } + t.file = f + t.path = path + return nil +} + +// Path returns the absolute spill path after any Write (may be empty if unused/failed). +func (t *Tee) Path() string { + if t == nil { + return "" + } + t.mu.Lock() + defer t.mu.Unlock() + return t.path +} + +// Close flushes and closes the spill file. +func (t *Tee) Close() error { + if t == nil { + return nil + } + t.mu.Lock() + defer t.mu.Unlock() + if t.file == nil { + return nil + } + err := t.file.Close() + t.file = nil + return err +} diff --git a/internal/tooloutput/spill_test.go b/internal/tooloutput/spill_test.go new file mode 100644 index 00000000..0b6c20a9 --- /dev/null +++ b/internal/tooloutput/spill_test.go @@ -0,0 +1,58 @@ +package tooloutput + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestBoundWithSpillWritesFullFile(t *testing.T) { + root := t.TempDir() + full := strings.Repeat("A", 2000) + "TAIL" + out := BoundWithSpill(full, 512, SpillOpts{ + RootDir: root, + ConversationID: "conv-1", + ExecutionID: "exec-1", + }) + if len(out) > 512 { + t.Fatalf("bounded output exceeds max: %d", len(out)) + } + if !strings.Contains(out, "") { + t.Fatalf("expected persisted-output notice: %q", out) + } + path := filepath.Join(root, "conversations", "conv-1", "trunc", "exec-1") + abs, err := filepath.Abs(path) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(out, abs) { + t.Fatalf("expected absolute path %q in notice: %q", abs, out) + } + got, err := os.ReadFile(abs) + if err != nil { + t.Fatal(err) + } + if string(got) != full { + t.Fatalf("spilled content mismatch: got %d want %d", len(got), len(full)) + } +} + +func TestTeeThenFormatPersistedFromFile(t *testing.T) { + root := t.TempDir() + tee := NewTee(SpillOpts{RootDir: root, ConversationID: "c", ExecutionID: "e"}) + full := strings.Repeat("xy", 100) + if _, err := tee.Write([]byte(full)); err != nil { + t.Fatal(err) + } + if err := tee.Close(); err != nil { + t.Fatal(err) + } + notice := FormatPersistedFromFile(tee.Path(), len(full), 512) + if len(notice) > 512 { + t.Fatalf("notice too long: %d", len(notice)) + } + if !strings.Contains(notice, tee.Path()) { + t.Fatalf("missing path in notice: %q", notice) + } +}