diff --git a/internal/agentfinalizer/decision.go b/internal/agentfinalizer/decision.go new file mode 100644 index 00000000..5aa0d00c --- /dev/null +++ b/internal/agentfinalizer/decision.go @@ -0,0 +1,266 @@ +package agentfinalizer + +import ( + "strings" + + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + "cyberstrike-ai/internal/multiagent" +) + +const ( + StatusCompleted = "completed" + StatusInProgress = "in_progress" + StatusBlocked = "blocked" + StatusFailed = "failed" + StatusCancelled = "cancelled" + StatusAwaitingHITL = "awaiting_hitl" + + ReasonVerified = "verified" + ReasonPendingTools = "pending_tool_executions" + ReasonEmptyResponse = "empty_response" + ReasonAwaitingHITL = "awaiting_hitl" + ReasonFailed = "failed" + ReasonCancelled = "cancelled" + ReasonMissingEvidence = "missing_execution_evidence" +) + +// Decision is the single contract that may promote an agent run to a final +// user-facing answer. Natural-language assistant text is only a candidate until +// this object says Finalizable. +type Decision struct { + Status string `json:"status"` + Finalizable bool `json:"finalizable"` + Finalized bool `json:"finalized"` + CompletionReason string `json:"completionReason"` + FinalText string `json:"finalText,omitempty"` + EvidenceVerified bool `json:"evidenceVerified"` + EvidenceRefs []string `json:"evidenceRefs,omitempty"` + PendingExecutionIDs []string `json:"pendingExecutionIds,omitempty"` + PendingToolRuns []string `json:"pendingToolRuns,omitempty"` + MissingChecks []string `json:"missingChecks,omitempty"` + AgentMode string `json:"agentMode,omitempty"` + ConversationID string `json:"conversationId,omitempty"` + AssistantMessageID string `json:"messageId,omitempty"` + CandidateResponseLen int `json:"candidateResponseLen,omitempty"` +} + +type Input struct { + Response string + MCPExecutionIDs []string + ConversationID string + AssistantMessageID string + AgentMode string + Status string + CompletionReason string + AwaitingHITL bool + RequireExecutionEvidence bool +} + +func FromRunResult(db *database.DB, result *multiagent.RunResult, in Input) Decision { + if result != nil { + if strings.TrimSpace(in.Response) == "" { + in.Response = result.Response + } + if len(in.MCPExecutionIDs) == 0 { + in.MCPExecutionIDs = result.MCPExecutionIDs + } + if strings.TrimSpace(in.Status) == "" { + in.Status = result.Status + } + if strings.TrimSpace(in.CompletionReason) == "" { + in.CompletionReason = result.CompletionReason + } + } + d := Decide(db, in) + if result != nil { + result.Finalized = d.Finalized + result.Status = d.Status + result.CompletionReason = d.CompletionReason + result.EvidenceVerified = d.EvidenceVerified + result.EvidenceRefs = append([]string(nil), d.EvidenceRefs...) + result.PendingExecutionIDs = append([]string(nil), d.PendingExecutionIDs...) + result.MissingChecks = append([]string(nil), d.MissingChecks...) + } + return d +} + +func Decide(db *database.DB, in Input) Decision { + text := strings.TrimSpace(in.Response) + status := strings.TrimSpace(in.Status) + if status == "" { + status = StatusCompleted + } + reason := strings.TrimSpace(in.CompletionReason) + if reason == "" { + reason = ReasonVerified + } + d := Decision{ + Status: status, + CompletionReason: reason, + FinalText: text, + EvidenceVerified: true, + EvidenceRefs: evidenceRefs(in.MCPExecutionIDs), + AgentMode: strings.TrimSpace(in.AgentMode), + ConversationID: strings.TrimSpace(in.ConversationID), + AssistantMessageID: strings.TrimSpace(in.AssistantMessageID), + CandidateResponseLen: len([]rune(text)), + } + + if in.AwaitingHITL { + d.Status = StatusAwaitingHITL + d.CompletionReason = ReasonAwaitingHITL + d.EvidenceVerified = false + d.MissingChecks = append(d.MissingChecks, "workflow is awaiting HITL approval") + return d + } + if isEmptyCandidate(text) { + d.Status = StatusBlocked + d.CompletionReason = ReasonEmptyResponse + d.EvidenceVerified = false + d.MissingChecks = append(d.MissingChecks, "assistant final text is empty or only an empty-response placeholder") + return d + } + switch status { + case StatusInProgress, StatusBlocked, StatusFailed, StatusCancelled, StatusAwaitingHITL: + d.Status = status + d.EvidenceVerified = false + if d.CompletionReason == ReasonVerified { + d.CompletionReason = status + } + d.MissingChecks = append(d.MissingChecks, "agent run status is "+status) + return d + } + + pending := pendingExecutions(db, in.MCPExecutionIDs) + if len(pending) > 0 { + d.Status = StatusInProgress + d.CompletionReason = ReasonPendingTools + d.EvidenceVerified = false + d.PendingExecutionIDs = pending + d.PendingToolRuns = append([]string(nil), pending...) + d.MissingChecks = append(d.MissingChecks, "tool execution still queued or running") + return d + } + + if in.RequireExecutionEvidence && !hasCompletedEvidence(db, in.MCPExecutionIDs) { + d.Status = StatusBlocked + d.CompletionReason = ReasonMissingEvidence + d.EvidenceVerified = false + d.MissingChecks = append(d.MissingChecks, "execution evidence is required but no completed tool execution was recorded") + return d + } + + d.Finalizable = true + d.Finalized = true + d.Status = StatusCompleted + if d.CompletionReason == "" { + d.CompletionReason = ReasonVerified + } + return d +} + +func ResponsePayload(d Decision, extra map[string]interface{}) map[string]interface{} { + out := map[string]interface{}{ + "finalized": d.Finalized, + "finalizable": d.Finalizable, + "status": d.Status, + "completionReason": d.CompletionReason, + "evidenceVerified": d.EvidenceVerified, + "evidenceRefs": d.EvidenceRefs, + "pendingExecutionIds": d.PendingExecutionIDs, + "pendingToolRuns": d.PendingToolRuns, + "missingChecks": d.MissingChecks, + } + if d.ConversationID != "" { + out["conversationId"] = d.ConversationID + } + if d.AssistantMessageID != "" { + out["messageId"] = d.AssistantMessageID + } + if d.AgentMode != "" { + out["agentMode"] = d.AgentMode + } + for k, v := range extra { + out[k] = v + } + return out +} + +func isEmptyCandidate(s string) bool { + s = strings.TrimSpace(s) + if s == "" { + return true + } + return strings.Contains(s, "no assistant text was captured") || + strings.Contains(s, "未捕获到助手文本输出") +} + +func evidenceRefs(ids []string) []string { + out := make([]string, 0, len(ids)) + seen := make(map[string]struct{}, len(ids)) + for _, id := range ids { + id = strings.TrimSpace(id) + if id == "" { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + out = append(out, "mcp_execution:"+id) + } + return out +} + +func pendingExecutions(db *database.DB, ids []string) []string { + if db == nil || len(ids) == 0 { + return nil + } + out := make([]string, 0) + seen := make(map[string]struct{}, len(ids)) + for _, id := range ids { + id = strings.TrimSpace(id) + if id == "" { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + exec, err := db.GetToolExecution(id) + if err != nil || exec == nil { + continue + } + switch strings.TrimSpace(exec.Status) { + case mcp.ToolExecutionStatusQueued, mcp.ToolExecutionStatusRunning: + out = append(out, id) + } + } + return out +} + +func hasCompletedEvidence(db *database.DB, ids []string) bool { + if db == nil || len(ids) == 0 { + return false + } + seen := make(map[string]struct{}, len(ids)) + for _, id := range ids { + id = strings.TrimSpace(id) + if id == "" { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + exec, err := db.GetToolExecution(id) + if err != nil || exec == nil { + continue + } + if strings.TrimSpace(exec.Status) == mcp.ToolExecutionStatusCompleted { + return true + } + } + return false +} diff --git a/internal/agentfinalizer/decision_test.go b/internal/agentfinalizer/decision_test.go new file mode 100644 index 00000000..c31755a3 --- /dev/null +++ b/internal/agentfinalizer/decision_test.go @@ -0,0 +1,132 @@ +package agentfinalizer + +import ( + "path/filepath" + "testing" + "time" + + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + + "go.uber.org/zap" +) + +func newDecisionTestDB(t *testing.T) *database.DB { + t.Helper() + db, err := database.NewDB(filepath.Join(t.TempDir(), "finalizer.db"), zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + t.Cleanup(func() { _ = db.Close() }) + return db +} + +func saveDecisionTestExecution(t *testing.T, db *database.DB, id, status string) { + t.Helper() + if err := db.SaveToolExecution(&mcp.ToolExecution{ + ID: id, + ToolName: "test::tool", + Arguments: map[string]interface{}{"input": id}, + Status: status, + StartTime: time.Now(), + }); err != nil { + t.Fatalf("SaveToolExecution(%s): %v", id, err) + } +} + +func TestDecideBlocksPendingToolExecutions(t *testing.T) { + db := newDecisionTestDB(t) + saveDecisionTestExecution(t, db, "run-queued", mcp.ToolExecutionStatusQueued) + saveDecisionTestExecution(t, db, "run-running", mcp.ToolExecutionStatusRunning) + saveDecisionTestExecution(t, db, "run-completed", mcp.ToolExecutionStatusCompleted) + + d := Decide(db, Input{ + Response: "工具还没全部结束时,这只是一段候选输出。", + MCPExecutionIDs: []string{"run-queued", "run-running", "run-completed"}, + }) + + if d.Finalizable || d.Finalized { + t.Fatalf("pending tools should not be finalizable: %+v", d) + } + if d.Status != StatusInProgress || d.CompletionReason != ReasonPendingTools { + t.Fatalf("status/reason = %s/%s, want %s/%s", d.Status, d.CompletionReason, StatusInProgress, ReasonPendingTools) + } + if got, want := len(d.PendingExecutionIDs), 2; got != want { + t.Fatalf("pending execution count = %d, want %d (%v)", got, want, d.PendingExecutionIDs) + } +} + +func TestDecideBlocksAwaitingHITLAndEmptyCandidate(t *testing.T) { + hitl := Decide(nil, Input{Response: "等待人工审批", AwaitingHITL: true}) + if hitl.Finalizable || hitl.Status != StatusAwaitingHITL || hitl.CompletionReason != ReasonAwaitingHITL { + t.Fatalf("HITL decision mismatch: %+v", hitl) + } + + empty := Decide(nil, Input{Response: "⚠️ Eino 执行完成,但未捕获到助手文本输出。"}) + if empty.Finalizable || empty.Status != StatusBlocked || empty.CompletionReason != ReasonEmptyResponse { + t.Fatalf("empty candidate decision mismatch: %+v", empty) + } +} + +func TestDecideBlocksWhenExecutionEvidenceIsRequiredButMissing(t *testing.T) { + d := Decide(nil, Input{ + Response: "任务已处理完成。", + RequireExecutionEvidence: true, + }) + if d.Finalizable || d.Finalized { + t.Fatalf("missing required execution evidence should not finalize: %+v", d) + } + if d.Status != StatusBlocked || d.CompletionReason != ReasonMissingEvidence { + t.Fatalf("status/reason = %s/%s, want %s/%s", d.Status, d.CompletionReason, StatusBlocked, ReasonMissingEvidence) + } + if d.EvidenceVerified { + t.Fatalf("missing required execution evidence should be marked unverified: %+v", d) + } + if len(d.MissingChecks) == 0 { + t.Fatalf("missing checks should explain the evidence gap: %+v", d) + } +} + +func TestDecideBlocksWhenOnlyFailedEvidenceIsRecorded(t *testing.T) { + db := newDecisionTestDB(t) + saveDecisionTestExecution(t, db, "run-failed", mcp.ToolExecutionStatusFailed) + saveDecisionTestExecution(t, db, "run-cancelled", mcp.ToolExecutionStatusCancelled) + + d := Decide(db, Input{ + Response: "任务已处理完成。", + MCPExecutionIDs: []string{"run-failed", "run-cancelled"}, + RequireExecutionEvidence: true, + }) + + if d.Finalizable || d.Finalized { + t.Fatalf("failed evidence should not satisfy required execution evidence: %+v", d) + } + if d.Status != StatusBlocked || d.CompletionReason != ReasonMissingEvidence { + t.Fatalf("status/reason = %s/%s, want %s/%s", d.Status, d.CompletionReason, StatusBlocked, ReasonMissingEvidence) + } +} + +func TestDecideFinalizesCompletedEvidence(t *testing.T) { + db := newDecisionTestDB(t) + saveDecisionTestExecution(t, db, "run-ok", mcp.ToolExecutionStatusCompleted) + + d := Decide(db, Input{ + Response: "任务已处理完成,见工具执行记录。", + MCPExecutionIDs: []string{"run-ok"}, + RequireExecutionEvidence: true, + }) + + if !d.Finalizable || !d.Finalized || d.Status != StatusCompleted { + t.Fatalf("completed execution should finalize: %+v", d) + } + if !d.EvidenceVerified || len(d.EvidenceRefs) != 1 { + t.Fatalf("evidence refs mismatch: %+v", d) + } +} + +func TestDecideAllowsInformationalAnswerWhenExecutionEvidenceIsNotRequired(t *testing.T) { + d := Decide(nil, Input{Response: "这是一个概念解释,不需要执行工具。"}) + if !d.Finalizable || !d.Finalized || d.Status != StatusCompleted { + t.Fatalf("informational response should finalize when execution evidence is not required: %+v", d) + } +} 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..f2583359 --- /dev/null +++ b/internal/database/tool_execution_args_lookup.go @@ -0,0 +1,53 @@ +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 + } + start := at.Add(-window) + end := at.Add(window) + rows, err := db.Query(` +SELECT id, arguments +FROM tool_executions +WHERE conversation_id = ? + AND tool_name = ? + AND julianday(start_time) BETWEEN julianday(?) AND julianday(?) +ORDER BY ABS(julianday(start_time) - julianday(?)) ASC, start_time ASC +LIMIT 1`, conversationID, toolName, 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 && r.logger != nil { + r.logger.Info("启动时已收尾孤儿 running 工具执行记录", zap.Int64("count", n)) + } +} + +func (r *ExecutionReconciler) activeExecutionIDs() map[string]struct{} { + ids := make(map[string]struct{}) + if r.mcpServer != nil { + for id := range r.mcpServer.ActiveRunningExecutionIDs() { + ids[id] = struct{}{} + } + } + if r.externalMgr != nil { + for id := range r.externalMgr.ActiveRunningExecutionIDs() { + ids[id] = struct{}{} + } + } + return ids +} + +// ReconcileStaleRunning finalizes running rows that are not tracked in-memory and older than staleRunningMinAge. +func (r *ExecutionReconciler) ReconcileStaleRunning() { + if r == nil || r.db == nil { + return + } + now := time.Now() + n, err := r.db.FinalizeStaleRunningToolExecutions(now, staleRunningMinAge, r.activeExecutionIDs(), "执行已中断(会话已结束)") + if err != nil { + if r.logger != nil { + r.logger.Warn("定期收尾 stale running 工具执行记录失败", zap.Error(err)) + } + return + } + if n > 0 && r.logger != nil { + r.logger.Info("已收尾 stale running 工具执行记录", zap.Int64("count", n)) + } +} + +// StartStaleRunningReconcileLoop periodically reconciles orphaned running tool executions. +func StartStaleRunningReconcileLoop(r *ExecutionReconciler, logger *zap.Logger) { + if r == nil { + return + } + go func() { + ticker := time.NewTicker(staleRunningReconcileGap) + defer ticker.Stop() + for range ticker.C { + r.ReconcileStaleRunning() + if logger != nil { + logger.Debug("monitor stale running reconcile tick completed") + } + } + }() +} diff --git a/internal/monitor/reconcile_test.go b/internal/monitor/reconcile_test.go new file mode 100644 index 00000000..0dfad0e9 --- /dev/null +++ b/internal/monitor/reconcile_test.go @@ -0,0 +1,38 @@ +package monitor + +import ( + "path/filepath" + "testing" + "time" + + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + + "go.uber.org/zap" +) + +func TestExecutionReconciler_ReconcileOnStartup(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "monitor.db") + db, err := database.NewDB(dbPath, zap.NewNop()) + if err != nil { + t.Fatalf("NewDB: %v", err) + } + defer db.Close() + + if err := db.SaveToolExecution(&mcp.ToolExecution{ + ID: "run-1", ToolName: "hydra", Status: "running", StartTime: time.Now().Add(-time.Hour), + }); err != nil { + t.Fatalf("SaveToolExecution: %v", err) + } + + r := NewExecutionReconciler(db, mcp.NewServer(zap.NewNop()), nil, zap.NewNop()) + r.ReconcileOnStartup() + + got, err := db.GetToolExecution("run-1") + if err != nil { + t.Fatalf("GetToolExecution: %v", err) + } + if got.Status != "orphaned" { + t.Fatalf("expected orphaned after startup reconcile, got %s", got.Status) + } +} diff --git a/internal/monitor/retention.go b/internal/monitor/retention.go new file mode 100644 index 00000000..d1ffb295 --- /dev/null +++ b/internal/monitor/retention.go @@ -0,0 +1,71 @@ +package monitor + +import ( + "time" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + + "go.uber.org/zap" +) + +const retentionPurgeInterval = time.Hour + +// Service manages MCP tool execution monitor retention. +type Service struct { + db *database.DB + cfg *config.Config + logger *zap.Logger +} + +// NewService creates a monitor retention service. +func NewService(db *database.DB, cfg *config.Config, logger *zap.Logger) *Service { + return &Service{db: db, cfg: cfg, logger: logger} +} + +// RetentionDays returns configured retention; 0 means keep forever. +func (s *Service) RetentionDays() int { + if s == nil || s.cfg == nil { + return config.MonitorConfig{}.RetentionDaysEffective() + } + return s.cfg.Monitor.RetentionDaysEffective() +} + +// PurgeExpired deletes tool execution rows older than retention_days when configured. +func (s *Service) PurgeExpired() { + if s == nil || s.db == nil || s.cfg == nil { + return + } + days := s.cfg.Monitor.RetentionDaysEffective() + if days <= 0 { + return + } + cutoff := time.Now().AddDate(0, 0, -days) + n, err := s.db.PurgeToolExecutionsBefore(cutoff) + if err != nil { + if s.logger != nil { + s.logger.Warn("清理过期 MCP 执行记录失败", zap.Error(err)) + } + return + } + if n > 0 && s.logger != nil { + s.logger.Info("已清理过期 MCP 执行记录", zap.Int64("deleted", n), zap.Int("retention_days", days)) + } +} + +// StartRetentionLoop periodically purges expired tool execution rows. +func StartRetentionLoop(s *Service, logger *zap.Logger) { + if s == nil { + return + } + go func() { + ticker := time.NewTicker(retentionPurgeInterval) + defer ticker.Stop() + for range ticker.C { + s.PurgeExpired() + if logger != nil { + logger.Debug("monitor retention tick completed") + } + } + }() +} diff --git a/internal/monitor/retention_test.go b/internal/monitor/retention_test.go new file mode 100644 index 00000000..40425fd6 --- /dev/null +++ b/internal/monitor/retention_test.go @@ -0,0 +1,94 @@ +package monitor + +import ( + "path/filepath" + "testing" + "time" + + "cyberstrike-ai/internal/config" + "cyberstrike-ai/internal/database" + "cyberstrike-ai/internal/mcp" + + "go.uber.org/zap" +) + +func TestServicePurgeExpired_respectsZeroRetention(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "monitor.db") + db, err := database.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: mustParseTime(t, "2020-01-01T00:00:00Z"), + } + if err := db.SaveToolExecution(exec); err != nil { + t.Fatalf("SaveToolExecution: %v", err) + } + + zero := 0 + svc := NewService(db, &config.Config{ + Monitor: config.MonitorConfig{RetentionDays: &zero}, + }, zap.NewNop()) + svc.PurgeExpired() + + if _, err := db.GetToolExecution("ancient"); err != nil { + t.Fatalf("record should remain when retention_days=0: %v", err) + } +} + +func TestServicePurgeExpired_deletesOldRows(t *testing.T) { + dbPath := filepath.Join(t.TempDir(), "monitor.db") + db, err := database.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: mustParseTime(t, "2020-01-01T00:00:00Z"), + } + if err := db.SaveToolExecution(exec); err != nil { + t.Fatalf("SaveToolExecution: %v", err) + } + + days := 90 + svc := NewService(db, &config.Config{ + Monitor: config.MonitorConfig{RetentionDays: &days}, + }, zap.NewNop()) + svc.PurgeExpired() + + if _, err := db.GetToolExecution("ancient"); err == nil { + t.Fatal("record should be purged when older than retention_days") + } +} + +func TestRetentionDaysEffective_defaults(t *testing.T) { + got := config.MonitorConfig{}.RetentionDaysEffective() + if got != 90 { + t.Fatalf("default = %d, want 90", got) + } + zero := 0 + cfg := config.MonitorConfig{RetentionDays: &zero} + if cfg.RetentionDaysEffective() != 0 { + t.Fatalf("zero = %d, want 0", cfg.RetentionDaysEffective()) + } +} + +func mustParseTime(t *testing.T, value string) time.Time { + t.Helper() + parsed, err := time.Parse(time.RFC3339, value) + if err != nil { + t.Fatalf("parse time: %v", err) + } + return parsed +} diff --git a/internal/skillpackage/content.go b/internal/skillpackage/content.go new file mode 100644 index 00000000..91a02310 --- /dev/null +++ b/internal/skillpackage/content.go @@ -0,0 +1,164 @@ +package skillpackage + +import ( + "fmt" + "regexp" + "strings" +) + +var reH2 = regexp.MustCompile(`(?m)^##\s+(.+)$`) + +const summaryContentRunes = 6000 + +type markdownSection struct { + Heading string + Title string + Content string +} + +func splitMarkdownSections(body string) []markdownSection { + body = strings.TrimSpace(body) + if body == "" { + return nil + } + idxs := reH2.FindAllStringIndex(body, -1) + titles := reH2.FindAllStringSubmatch(body, -1) + if len(idxs) == 0 { + return []markdownSection{{ + Heading: "", + Title: "_body", + Content: body, + }} + } + var out []markdownSection + for i := range idxs { + title := strings.TrimSpace(titles[i][1]) + start := idxs[i][0] + end := len(body) + if i+1 < len(idxs) { + end = idxs[i+1][0] + } + chunk := strings.TrimSpace(body[start:end]) + out = append(out, markdownSection{ + Heading: "## " + title, + Title: title, + Content: chunk, + }) + } + return out +} + +func deriveSections(body string) []SkillSection { + md := splitMarkdownSections(body) + out := make([]SkillSection, 0, len(md)) + for _, ms := range md { + if ms.Title == "_body" { + continue + } + out = append(out, SkillSection{ + ID: slugifySectionID(ms.Title), + Title: ms.Title, + Heading: ms.Heading, + Level: 2, + }) + } + return out +} + +func slugifySectionID(title string) string { + title = strings.TrimSpace(strings.ToLower(title)) + if title == "" { + return "section" + } + var b strings.Builder + for _, r := range title { + switch { + case r >= 'a' && r <= 'z', r >= '0' && r <= '9': + b.WriteRune(r) + case r == ' ', r == '-', r == '_': + b.WriteRune('-') + } + } + s := strings.Trim(b.String(), "-") + if s == "" { + return "section" + } + return s +} + +func findSectionContent(sections []markdownSection, sec string) string { + sec = strings.TrimSpace(sec) + if sec == "" { + return "" + } + want := strings.ToLower(sec) + for _, s := range sections { + if strings.EqualFold(slugifySectionID(s.Title), want) || strings.EqualFold(s.Title, sec) { + return s.Content + } + if strings.EqualFold(strings.ReplaceAll(s.Title, " ", "-"), want) { + return s.Content + } + } + return "" +} + +func buildSummaryMarkdown(name, description string, tags []string, scripts []SkillScriptInfo, sections []SkillSection, body string) string { + var b strings.Builder + if description != "" { + b.WriteString(description) + b.WriteString("\n\n") + } + if len(tags) > 0 { + b.WriteString("**Tags**: ") + b.WriteString(strings.Join(tags, ", ")) + b.WriteString("\n\n") + } + if len(scripts) > 0 { + b.WriteString("### Bundled scripts\n\n") + for _, sc := range scripts { + line := "- `" + sc.RelPath + "`" + if sc.Description != "" { + line += " — " + sc.Description + } + b.WriteString(line) + b.WriteString("\n") + } + b.WriteString("\n") + } + if len(sections) > 0 { + b.WriteString("### Sections\n\n") + for _, sec := range sections { + line := "- **" + sec.ID + "**" + if sec.Title != "" && sec.Title != sec.ID { + line += ": " + sec.Title + } + b.WriteString(line) + b.WriteString("\n") + } + b.WriteString("\n") + } + mdSecs := splitMarkdownSections(body) + preview := body + if len(mdSecs) > 0 && mdSecs[0].Title != "_body" { + preview = mdSecs[0].Content + } + b.WriteString("### Preview (SKILL.md)\n\n") + b.WriteString(truncateRunes(strings.TrimSpace(preview), summaryContentRunes)) + b.WriteString("\n\n---\n\n_(Summary for admin UI. Agents use Eino `skill` tool for full SKILL.md progressive loading.)_") + if name != "" { + b.WriteString(fmt.Sprintf("\n\n_Skill name: %s_", name)) + } + return b.String() +} + +func truncateRunes(s string, max int) string { + if max <= 0 || s == "" { + return s + } + r := []rune(s) + if len(r) <= max { + return s + } + return string(r[:max]) + "…" +} diff --git a/internal/skillpackage/frontmatter.go b/internal/skillpackage/frontmatter.go new file mode 100644 index 00000000..905156b1 --- /dev/null +++ b/internal/skillpackage/frontmatter.go @@ -0,0 +1,114 @@ +package skillpackage + +import ( + "fmt" + "strings" + + "gopkg.in/yaml.v3" +) + +// ExtractSkillMDFrontMatterYAML returns the YAML source inside the first --- ... --- block and the markdown body. +func ExtractSkillMDFrontMatterYAML(raw []byte) (fmYAML string, body string, err error) { + text := strings.TrimPrefix(string(raw), "\ufeff") + if strings.TrimSpace(text) == "" { + return "", "", fmt.Errorf("SKILL.md is empty") + } + lines := strings.Split(text, "\n") + if len(lines) < 2 || strings.TrimSpace(lines[0]) != "---" { + return "", "", fmt.Errorf("SKILL.md must start with YAML front matter (---) per Agent Skills standard") + } + var fmLines []string + i := 1 + for i < len(lines) { + if strings.TrimSpace(lines[i]) == "---" { + break + } + fmLines = append(fmLines, lines[i]) + i++ + } + if i >= len(lines) { + return "", "", fmt.Errorf("SKILL.md: front matter must end with a line containing only ---") + } + body = strings.Join(lines[i+1:], "\n") + body = strings.TrimSpace(body) + fmYAML = strings.Join(fmLines, "\n") + return fmYAML, body, nil +} + +// ParseSkillMD parses SKILL.md YAML head + body. +func ParseSkillMD(raw []byte) (*SkillManifest, string, error) { + fmYAML, body, err := ExtractSkillMDFrontMatterYAML(raw) + if err != nil { + return nil, "", err + } + var m SkillManifest + if err := yaml.Unmarshal([]byte(fmYAML), &m); err != nil { + return nil, "", fmt.Errorf("SKILL.md front matter: %w", err) + } + return &m, body, nil +} + +type skillFrontMatterExport struct { + Name string `yaml:"name"` + Description string `yaml:"description"` + License string `yaml:"license,omitempty"` + Compatibility string `yaml:"compatibility,omitempty"` + Metadata map[string]any `yaml:"metadata,omitempty"` + AllowedTools string `yaml:"allowed-tools,omitempty"` +} + +// BuildSkillMD serializes SKILL.md per agentskills.io. +func BuildSkillMD(m *SkillManifest, body string) ([]byte, error) { + if m == nil { + return nil, fmt.Errorf("nil manifest") + } + fm := skillFrontMatterExport{ + Name: strings.TrimSpace(m.Name), + Description: strings.TrimSpace(m.Description), + License: strings.TrimSpace(m.License), + Compatibility: strings.TrimSpace(m.Compatibility), + AllowedTools: strings.TrimSpace(m.AllowedTools), + } + if len(m.Metadata) > 0 { + fm.Metadata = m.Metadata + } + head, err := yaml.Marshal(&fm) + if err != nil { + return nil, err + } + s := strings.TrimSpace(string(head)) + out := "---\n" + s + "\n---\n\n" + strings.TrimSpace(body) + "\n" + return []byte(out), nil +} + +func manifestTags(m *SkillManifest) []string { + if m == nil || m.Metadata == nil { + return nil + } + var out []string + if raw, ok := m.Metadata["tags"]; ok { + switch v := raw.(type) { + case []any: + for _, x := range v { + if s, ok := x.(string); ok && s != "" { + out = append(out, s) + } + } + case []string: + out = append(out, v...) + } + } + return out +} + +func versionFromMetadata(m *SkillManifest) string { + if m == nil || m.Metadata == nil { + return "" + } + if v, ok := m.Metadata["version"]; ok { + if s, ok := v.(string); ok { + return strings.TrimSpace(s) + } + } + return "" +} diff --git a/internal/skillpackage/io.go b/internal/skillpackage/io.go new file mode 100644 index 00000000..8a2b7222 --- /dev/null +++ b/internal/skillpackage/io.go @@ -0,0 +1,200 @@ +package skillpackage + +import ( + "fmt" + "io/fs" + "os" + "path/filepath" + "strings" +) + +const ( + maxPackageFiles = 4000 + maxPackageDepth = 24 + maxScriptsDepth = 24 + defaultMaxRead = 10 << 20 +) + +// SafeRelPath resolves rel inside root (no ..). +func SafeRelPath(root, rel string) (string, error) { + rel = strings.TrimSpace(rel) + rel = filepath.ToSlash(rel) + rel = strings.TrimPrefix(rel, "/") + if rel == "" || rel == "." { + return "", fmt.Errorf("empty resource path") + } + if strings.Contains(rel, "..") { + return "", fmt.Errorf("invalid path %q", rel) + } + abs := filepath.Join(root, filepath.FromSlash(rel)) + cleanRoot := filepath.Clean(root) + cleanAbs := filepath.Clean(abs) + relOut, err := filepath.Rel(cleanRoot, cleanAbs) + if err != nil || relOut == ".." || strings.HasPrefix(relOut, ".."+string(filepath.Separator)) { + return "", fmt.Errorf("path escapes skill directory: %q", rel) + } + return cleanAbs, nil +} + +// ListPackageFiles lists files under a skill directory. +func ListPackageFiles(skillsRoot, skillID string) ([]PackageFileInfo, error) { + root := SkillDir(skillsRoot, skillID) + if _, err := ResolveSKILLPath(root); err != nil { + return nil, fmt.Errorf("skill %q: %w", skillID, err) + } + var out []PackageFileInfo + err := filepath.WalkDir(root, func(path string, d fs.DirEntry, err error) error { + if err != nil { + return err + } + rel, e := filepath.Rel(root, path) + if e != nil { + return e + } + if rel == "." { + return nil + } + depth := strings.Count(rel, string(os.PathSeparator)) + if depth > maxPackageDepth { + if d.IsDir() { + return filepath.SkipDir + } + return nil + } + if strings.HasPrefix(d.Name(), ".") { + if d.IsDir() { + return filepath.SkipDir + } + return nil + } + if len(out) >= maxPackageFiles { + return fmt.Errorf("skill package exceeds %d files", maxPackageFiles) + } + fi, err := d.Info() + if err != nil { + return err + } + out = append(out, PackageFileInfo{ + Path: filepath.ToSlash(rel), + Size: fi.Size(), + IsDir: d.IsDir(), + }) + return nil + }) + return out, err +} + +// ReadPackageFile reads a file relative to the skill package. +func ReadPackageFile(skillsRoot, skillID, relPath string, maxBytes int64) ([]byte, error) { + if maxBytes <= 0 { + maxBytes = defaultMaxRead + } + root := SkillDir(skillsRoot, skillID) + abs, err := SafeRelPath(root, relPath) + if err != nil { + return nil, err + } + fi, err := os.Stat(abs) + if err != nil { + return nil, err + } + if fi.IsDir() { + return nil, fmt.Errorf("path is a directory") + } + if fi.Size() > maxBytes { + return readFileHead(abs, maxBytes) + } + return os.ReadFile(abs) +} + +// WritePackageFile writes a file inside the skill package. +func WritePackageFile(skillsRoot, skillID, relPath string, content []byte) error { + root := SkillDir(skillsRoot, skillID) + if _, err := ResolveSKILLPath(root); err != nil { + return fmt.Errorf("skill %q: %w", skillID, err) + } + abs, err := SafeRelPath(root, relPath) + if err != nil { + return err + } + if err := os.MkdirAll(filepath.Dir(abs), 0755); err != nil { + return err + } + return os.WriteFile(abs, content, 0644) +} + +func readFileHead(path string, max int64) ([]byte, error) { + f, err := os.Open(path) + if err != nil { + return nil, err + } + defer f.Close() + buf := make([]byte, max) + n, err := f.Read(buf) + if err != nil && n == 0 { + return nil, err + } + return buf[:n], nil +} + +func listScripts(skillsRoot, skillID string) ([]SkillScriptInfo, error) { + root := filepath.Join(SkillDir(skillsRoot, skillID), "scripts") + st, err := os.Stat(root) + if err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, err + } + if !st.IsDir() { + return nil, nil + } + var out []SkillScriptInfo + err = filepath.WalkDir(root, func(path string, d os.DirEntry, err error) error { + if err != nil { + return err + } + rel, e := filepath.Rel(root, path) + if e != nil { + return e + } + if rel == "." { + return nil + } + if d.IsDir() { + if strings.HasPrefix(d.Name(), ".") { + return filepath.SkipDir + } + if strings.Count(rel, string(os.PathSeparator)) >= maxScriptsDepth { + return filepath.SkipDir + } + return nil + } + if strings.HasPrefix(d.Name(), ".") { + return nil + } + relSkill := filepath.Join("scripts", rel) + full := filepath.Join(root, rel) + fi, err := os.Stat(full) + if err != nil || fi.IsDir() { + return nil + } + out = append(out, SkillScriptInfo{ + Name: filepath.Base(rel), + RelPath: filepath.ToSlash(relSkill), + Size: fi.Size(), + }) + return nil + }) + return out, err +} + +func countNonDirFiles(files []PackageFileInfo) int { + n := 0 + for _, f := range files { + if !f.IsDir && f.Path != "SKILL.md" { + n++ + } + } + return n +} diff --git a/internal/skillpackage/layout.go b/internal/skillpackage/layout.go new file mode 100644 index 00000000..275e1924 --- /dev/null +++ b/internal/skillpackage/layout.go @@ -0,0 +1,66 @@ +package skillpackage + +import ( + "fmt" + "os" + "path/filepath" + "strings" +) + +// SkillDir returns the absolute path to a skill package directory. +func SkillDir(skillsRoot, skillID string) string { + return filepath.Join(skillsRoot, skillID) +} + +// ResolveSKILLPath returns SKILL.md path or error if missing. +func ResolveSKILLPath(skillPath string) (string, error) { + md := filepath.Join(skillPath, "SKILL.md") + if st, err := os.Stat(md); err != nil || st.IsDir() { + return "", fmt.Errorf("missing SKILL.md in %q (Agent Skills standard)", filepath.Base(skillPath)) + } + return md, nil +} + +// SkillsRootFromConfig resolves cfg.SkillsDir relative to the config file directory. +func SkillsRootFromConfig(skillsDir string, configPath string) string { + if skillsDir == "" { + skillsDir = "skills" + } + configDir := filepath.Dir(configPath) + if !filepath.IsAbs(skillsDir) { + skillsDir = filepath.Join(configDir, skillsDir) + } + return skillsDir +} + +// DirLister lists skill package directory names under SkillsRoot. +type DirLister struct { + SkillsRoot string +} + +// ListSkills returns skill package directory names that contain SKILL.md. +func (d DirLister) ListSkills() ([]string, error) { + return ListSkillDirNames(d.SkillsRoot) +} + +// ListSkillDirNames returns subdirectory names under skillsRoot that contain SKILL.md. +func ListSkillDirNames(skillsRoot string) ([]string, error) { + if _, err := os.Stat(skillsRoot); os.IsNotExist(err) { + return nil, nil + } + entries, err := os.ReadDir(skillsRoot) + if err != nil { + return nil, fmt.Errorf("read skills directory: %w", err) + } + var names []string + for _, entry := range entries { + if !entry.IsDir() || strings.HasPrefix(entry.Name(), ".") { + continue + } + skillPath := filepath.Join(skillsRoot, entry.Name()) + if _, err := ResolveSKILLPath(skillPath); err == nil { + names = append(names, entry.Name()) + } + } + return names, nil +} diff --git a/internal/skillpackage/service.go b/internal/skillpackage/service.go new file mode 100644 index 00000000..52dbe90a --- /dev/null +++ b/internal/skillpackage/service.go @@ -0,0 +1,155 @@ +package skillpackage + +import ( + "fmt" + "os" + "sort" + "strings" +) + +// ListSkillSummaries scans skillsRoot and returns index rows for the admin API. +func ListSkillSummaries(skillsRoot string) ([]SkillSummary, error) { + names, err := ListSkillDirNames(skillsRoot) + if err != nil { + return nil, err + } + sort.Strings(names) + out := make([]SkillSummary, 0, len(names)) + for _, dirName := range names { + su, err := loadSummary(skillsRoot, dirName) + if err != nil { + continue + } + out = append(out, su) + } + return out, nil +} + +func loadSummary(skillsRoot, dirName string) (SkillSummary, error) { + skillPath := SkillDir(skillsRoot, dirName) + mdPath, err := ResolveSKILLPath(skillPath) + if err != nil { + return SkillSummary{}, err + } + raw, err := os.ReadFile(mdPath) + if err != nil { + return SkillSummary{}, err + } + man, _, err := ParseSkillMD(raw) + if err != nil { + return SkillSummary{}, err + } + if err := ValidateAgentSkillManifestInPackage(man, dirName); err != nil { + return SkillSummary{}, err + } + fi, err := os.Stat(mdPath) + if err != nil { + return SkillSummary{}, err + } + pfiles, err := ListPackageFiles(skillsRoot, dirName) + if err != nil { + return SkillSummary{}, err + } + nFiles := 0 + for _, p := range pfiles { + if !p.IsDir { + nFiles++ + } + } + scripts, err := listScripts(skillsRoot, dirName) + if err != nil { + return SkillSummary{}, err + } + ver := versionFromMetadata(man) + return SkillSummary{ + ID: dirName, + DirName: dirName, + Name: man.Name, + Description: man.Description, + Version: ver, + Path: skillPath, + Tags: manifestTags(man), + ScriptCount: len(scripts), + FileCount: nFiles, + FileSize: fi.Size(), + ModTime: fi.ModTime().Format("2006-01-02 15:04:05"), + Progressive: true, + }, nil +} + +// LoadOptions mirrors legacy API query params for the web admin. +type LoadOptions struct { + Depth string // summary | full + Section string +} + +// LoadSkill returns manifest + body + package listing for admin. +func LoadSkill(skillsRoot, skillID string, opt LoadOptions) (*SkillView, error) { + skillPath := SkillDir(skillsRoot, skillID) + mdPath, err := ResolveSKILLPath(skillPath) + if err != nil { + return nil, err + } + raw, err := os.ReadFile(mdPath) + if err != nil { + return nil, err + } + man, body, err := ParseSkillMD(raw) + if err != nil { + return nil, err + } + if err := ValidateAgentSkillManifestInPackage(man, skillID); err != nil { + return nil, err + } + pfiles, err := ListPackageFiles(skillsRoot, skillID) + if err != nil { + return nil, err + } + scripts, err := listScripts(skillsRoot, skillID) + if err != nil { + return nil, err + } + sort.Slice(scripts, func(i, j int) bool { return scripts[i].RelPath < scripts[j].RelPath }) + sections := deriveSections(body) + ver := versionFromMetadata(man) + v := &SkillView{ + DirName: skillID, + Name: man.Name, + Description: man.Description, + Content: body, + Path: skillPath, + Version: ver, + Tags: manifestTags(man), + Scripts: scripts, + Sections: sections, + PackageFiles: pfiles, + } + depth := strings.ToLower(strings.TrimSpace(opt.Depth)) + if depth == "" { + depth = "full" + } + sec := strings.TrimSpace(opt.Section) + if sec != "" { + mds := splitMarkdownSections(body) + chunk := findSectionContent(mds, sec) + if chunk == "" { + v.Content = fmt.Sprintf("_(section %q not found in SKILL.md for skill %s)_", sec, skillID) + } else { + v.Content = chunk + } + return v, nil + } + if depth == "summary" { + v.Content = buildSummaryMarkdown(man.Name, man.Description, v.Tags, scripts, sections, body) + } + return v, nil +} + +// ReadScriptText returns file content as string (for HTTP resource_path). +func ReadScriptText(skillsRoot, skillID, relPath string, maxBytes int64) (string, error) { + b, err := ReadPackageFile(skillsRoot, skillID, relPath, maxBytes) + if err != nil { + return "", err + } + return string(b), nil +} diff --git a/internal/skillpackage/types.go b/internal/skillpackage/types.go new file mode 100644 index 00000000..bf313425 --- /dev/null +++ b/internal/skillpackage/types.go @@ -0,0 +1,67 @@ +// Package skillpackage provides filesystem-backed Agent Skills layout (SKILL.md + package files) +// for HTTP admin APIs. Runtime discovery and progressive loading for agents use Eino ADK skill middleware. +package skillpackage + +// SkillManifest is parsed from SKILL.md front matter (https://agentskills.io/specification.md). +type SkillManifest struct { + Name string `yaml:"name"` + Description string `yaml:"description"` + License string `yaml:"license,omitempty"` + Compatibility string `yaml:"compatibility,omitempty"` + Metadata map[string]any `yaml:"metadata,omitempty"` + AllowedTools string `yaml:"allowed-tools,omitempty"` +} + +// SkillSummary is API metadata for one skill directory. +type SkillSummary struct { + ID string `json:"id"` + DirName string `json:"dir_name"` + Name string `json:"name"` + Description string `json:"description"` + Version string `json:"version"` + Path string `json:"path"` + Tags []string `json:"tags"` + Triggers []string `json:"triggers,omitempty"` + ScriptCount int `json:"script_count"` + FileCount int `json:"file_count"` + FileSize int64 `json:"file_size"` + ModTime string `json:"mod_time"` + Progressive bool `json:"progressive"` +} + +// SkillScriptInfo describes a file under scripts/. +type SkillScriptInfo struct { + Name string `json:"name"` + RelPath string `json:"rel_path"` + Description string `json:"description,omitempty"` + Size int64 `json:"size"` +} + +// SkillSection is derived from ## headings in SKILL.md. +type SkillSection struct { + ID string `json:"id"` + Title string `json:"title"` + Heading string `json:"heading"` + Level int `json:"level"` +} + +// PackageFileInfo describes one file inside a package. +type PackageFileInfo struct { + Path string `json:"path"` + Size int64 `json:"size"` + IsDir bool `json:"is_dir,omitempty"` +} + +// SkillView is a loaded package for admin / API. +type SkillView struct { + DirName string `json:"dir_name"` + Name string `json:"name"` + Description string `json:"description"` + Content string `json:"content"` + Path string `json:"path"` + Version string `json:"version"` + Tags []string `json:"tags"` + Scripts []SkillScriptInfo `json:"scripts,omitempty"` + Sections []SkillSection `json:"sections,omitempty"` + PackageFiles []PackageFileInfo `json:"package_files,omitempty"` +} diff --git a/internal/skillpackage/validate.go b/internal/skillpackage/validate.go new file mode 100644 index 00000000..79d8255c --- /dev/null +++ b/internal/skillpackage/validate.go @@ -0,0 +1,102 @@ +package skillpackage + +import ( + "fmt" + "strings" + "unicode/utf8" + + "gopkg.in/yaml.v3" +) + +var agentSkillsSpecFrontMatterKeys = map[string]struct{}{ + "name": {}, "description": {}, "license": {}, "compatibility": {}, + "metadata": {}, "allowed-tools": {}, +} + +// ValidateAgentSkillManifest enforces Agent Skills rules for name and description. +func ValidateAgentSkillManifest(m *SkillManifest) error { + if m == nil { + return fmt.Errorf("skill manifest is nil") + } + if strings.TrimSpace(m.Name) == "" { + return fmt.Errorf("SKILL.md front matter: name is required") + } + if strings.TrimSpace(m.Description) == "" { + return fmt.Errorf("SKILL.md front matter: description is required") + } + if utf8.RuneCountInString(m.Name) > 64 { + return fmt.Errorf("name exceeds 64 characters (Agent Skills limit)") + } + if utf8.RuneCountInString(m.Description) > 1024 { + return fmt.Errorf("description exceeds 1024 characters (Agent Skills limit)") + } + if m.Name != strings.ToLower(m.Name) { + return fmt.Errorf("name must be lowercase (Agent Skills)") + } + for _, r := range m.Name { + if !((r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') || r == '-') { + return fmt.Errorf("name must contain only lowercase letters, numbers, hyphens (Agent Skills)") + } + } + if strings.HasPrefix(m.Name, "-") || strings.HasSuffix(m.Name, "-") { + return fmt.Errorf("name must not start or end with a hyphen (Agent Skills spec)") + } + if strings.Contains(m.Name, "--") { + return fmt.Errorf("name must not contain consecutive hyphens (Agent Skills spec)") + } + lname := strings.ToLower(m.Name) + if strings.Contains(lname, "anthropic") || strings.Contains(lname, "claude") { + return fmt.Errorf("name must not contain reserved words anthropic or claude") + } + return nil +} + +// ValidateAgentSkillManifestInPackage checks manifest and that name matches package directory. +func ValidateAgentSkillManifestInPackage(m *SkillManifest, packageDirName string) error { + if err := ValidateAgentSkillManifest(m); err != nil { + return err + } + if strings.TrimSpace(packageDirName) == "" { + return nil + } + if m.Name != packageDirName { + return fmt.Errorf("SKILL.md name %q must match directory name %q (Agent Skills spec)", m.Name, packageDirName) + } + return nil +} + +// ValidateOfficialFrontMatterTopLevelKeys rejects keys not in the open spec. +func ValidateOfficialFrontMatterTopLevelKeys(fmYAML string) error { + var top map[string]interface{} + if err := yaml.Unmarshal([]byte(fmYAML), &top); err != nil { + return fmt.Errorf("SKILL.md front matter: %w", err) + } + for k := range top { + if _, ok := agentSkillsSpecFrontMatterKeys[k]; !ok { + return fmt.Errorf("SKILL.md front matter: unsupported key %q (allowed: name, description, license, compatibility, metadata, allowed-tools — see https://agentskills.io/specification.md)", k) + } + } + return nil +} + +// ValidateSkillMDPackage validates SKILL.md bytes for writes. +func ValidateSkillMDPackage(raw []byte, packageDirName string) error { + fmYAML, body, err := ExtractSkillMDFrontMatterYAML(raw) + if err != nil { + return err + } + if err := ValidateOfficialFrontMatterTopLevelKeys(fmYAML); err != nil { + return err + } + if strings.TrimSpace(body) == "" { + return fmt.Errorf("SKILL.md: markdown body after front matter must not be empty") + } + var fm SkillManifest + if err := yaml.Unmarshal([]byte(fmYAML), &fm); err != nil { + return fmt.Errorf("SKILL.md front matter: %w", err) + } + if c := strings.TrimSpace(fm.Compatibility); c != "" && utf8.RuneCountInString(c) > 500 { + return fmt.Errorf("compatibility exceeds 500 characters (Agent Skills spec)") + } + return ValidateAgentSkillManifestInPackage(&fm, packageDirName) +}