From 84e99220ff4e323a83e39f6c39ea8ce2e969708c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=85=AC=E6=98=8E?= <83812544+Ed1s0nZ@users.noreply.github.com> Date: Fri, 31 Jul 2026 21:12:01 +0800 Subject: [PATCH] Add files via upload --- internal/database/asset.go | 158 +++++++++++++++++- internal/database/asset_test.go | 7 +- internal/database/conversation.go | 2 + internal/database/database.go | 8 +- .../database/process_details_summary_test.go | 8 +- internal/database/vulnerability.go | 31 ++++ 6 files changed, 203 insertions(+), 11 deletions(-) diff --git a/internal/database/asset.go b/internal/database/asset.go index b42184b0..846bb25a 100644 --- a/internal/database/asset.go +++ b/internal/database/asset.go @@ -13,6 +13,7 @@ import ( "unicode/utf8" "github.com/google/uuid" + "go.uber.org/zap" "golang.org/x/net/idna" ) @@ -50,6 +51,7 @@ type Asset struct { LastScanTaskID string `json:"last_scan_task_id,omitempty"` VulnerabilityCount int `json:"vulnerability_count"` RiskLevel string `json:"risk_level"` + RiskScore int `json:"-"` OwnerUserID string `json:"-"` } @@ -443,15 +445,15 @@ func assetWhere(filter AssetListFilter, access RBACListAccess) (string, []interf args = append(args, *filter.Port) } if filter.RiskLevel != "" { - query += " AND " + assetRiskLevelExpr + " = ?" + query += " AND " + assetRiskLevelCachedExpr + " = ?" args = append(args, strings.ToLower(strings.TrimSpace(filter.RiskLevel))) } if filter.MinVulnerabilities != nil { - query += " AND " + assetVulnerabilityCountExpr + " >= ?" + query += " AND " + assetVulnerabilityCountCachedExpr + " >= ?" args = append(args, *filter.MinVulnerabilities) } if filter.MaxVulnerabilities != nil { - query += " AND " + assetVulnerabilityCountExpr + " <= ?" + query += " AND " + assetVulnerabilityCountCachedExpr + " <= ?" args = append(args, *filter.MaxVulnerabilities) } for _, item := range []struct { @@ -573,20 +575,24 @@ const assetVulnerabilityMatchExpr = `( const assetVulnerabilityCountExpr = `(SELECT COUNT(DISTINCT v.id) FROM vulnerabilities v WHERE ` + assetVulnerabilityMatchExpr + `)` -const assetRiskScoreExpr = `COALESCE(( +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 assetRiskLevelExpr = `(CASE WHEN ` + assetEffectiveLastScanExpr + ` IS NULL THEN 'unassessed' ELSE CASE ` + assetRiskScoreExpr + ` +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,''), - ` + assetVulnerabilityCountExpr + `,` + assetRiskLevelExpr + ` + 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. @@ -601,6 +607,9 @@ func (db *DB) MarkAssetScanned(id, conversationID, queueID, taskID string, acces if n == 0 { return sql.ErrNoRows } + if err := db.RefreshAssetRiskCache(id); err != nil { + return err + } return nil } @@ -629,6 +638,9 @@ func (db *DB) CompleteAssetScan(id, conversationID string, access RBACListAccess if n == 0 { return sql.ErrNoRows } + if err := db.RefreshAssetRiskCache(id); err != nil { + return err + } return nil } @@ -638,6 +650,136 @@ func (db *DB) BatchTaskBelongsToQueue(taskID, queueID string) bool { 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 @@ -726,9 +868,9 @@ func assetOrderBy(sortBy, sortOrder string) string { case "port": expression = "assets.port" case "vulnerability_count": - expression = assetVulnerabilityCountExpr + expression = assetVulnerabilityCountCachedExpr case "risk_level": - expression = assetRiskScoreExpr + expression = assetRiskScoreCachedExpr default: expression = "assets.last_seen_at" } diff --git a/internal/database/asset_test.go b/internal/database/asset_test.go index c2929a0f..b9418d9b 100644 --- a/internal/database/asset_test.go +++ b/internal/database/asset_test.go @@ -383,7 +383,12 @@ func TestAssetScanLinkReturnsTimeAndRelatedVulnerabilities(t *testing.T) { if linked.LastScanAt == nil || linked.LastScanConversationID != conv.ID || linked.VulnerabilityCount != 1 || linked.RiskLevel != "high" { t.Fatalf("unexpected scan metadata: %#v", linked) } - if _, err := db.Exec(`UPDATE vulnerabilities SET status='fixed' WHERE conversation_id=?`, conv.ID); err != nil { + 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}) diff --git a/internal/database/conversation.go b/internal/database/conversation.go index f3ac24b3..e54cbeba 100644 --- a/internal/database/conversation.go +++ b/internal/database/conversation.go @@ -1424,6 +1424,7 @@ type ProcessDetailsSummary struct { 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"` @@ -1552,6 +1553,7 @@ func (db *DB) GetProcessDetailsSummary(messageID string) (*ProcessDetailsSummary if toolCallID != "" { lastMatchedToolIndexByCallID[toolCallID] = idx } + summary.ToolExecutions[idx].ResultDetailID = strings.TrimSpace(detailID) if summary.ToolExecutions[idx].ToolName == "" { summary.ToolExecutions[idx].ToolName = toolName } diff --git a/internal/database/database.go b/internal/database/database.go index 3c40b261..35884987 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -421,6 +421,7 @@ func (db *DB) initTables() error { 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 @@ -745,6 +746,9 @@ func (db *DB) initTables() error { 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); @@ -977,7 +981,6 @@ func (db *DB) initTables() error { if _, err := db.Exec(createIndexes); err != nil { return fmt.Errorf("创建索引失败: %w", err) } - db.logger.Debug("数据库表初始化完成") return nil } @@ -1027,6 +1030,9 @@ func (db *DB) migrateAssetsTable() error { {"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 diff --git a/internal/database/process_details_summary_test.go b/internal/database/process_details_summary_test.go index 39510da1..d436f4b0 100644 --- a/internal/database/process_details_summary_test.go +++ b/internal/database/process_details_summary_test.go @@ -22,10 +22,13 @@ func TestProcessDetailsSummaryDoesNotGuessIDLessResultsByOrder(t *testing.T) { {"toolName": "http-framework-test", "success": true}, {"toolName": "http-framework-test", "success": true}, } + var resultIDs []string for _, result := range results { - if err := db.AddProcessDetail(messageID, conversationID, "tool_result", "result", result); err != nil { + 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) @@ -39,6 +42,9 @@ func TestProcessDetailsSummaryDoesNotGuessIDLessResultsByOrder(t *testing.T) { 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" { diff --git a/internal/database/vulnerability.go b/internal/database/vulnerability.go index 664b2ba8..2eede0fa 100644 --- a/internal/database/vulnerability.go +++ b/internal/database/vulnerability.go @@ -191,6 +191,7 @@ func (db *DB) CreateVulnerability(vuln *Vulnerability) (*Vulnerability, error) { if err != nil { return nil, fmt.Errorf("创建漏洞失败: %w", err) } + db.refreshAssetRiskCacheForConversationsBestEffort(vuln.ConversationID) return vuln, nil } @@ -299,6 +300,8 @@ func (db *DB) CountVulnerabilitiesForAccess(filter VulnerabilityListFilter, acce // 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 @@ -318,6 +321,7 @@ func (db *DB) UpdateVulnerability(id string, vuln *Vulnerability) error { return fmt.Errorf("更新漏洞失败: %w", err) } + db.refreshAssetRiskCacheForConversationsBestEffort(oldConversationID, vuln.ConversationID) return nil } @@ -337,6 +341,10 @@ func (db *DB) DeleteVulnerabilitiesByFilterForAccess(filter VulnerabilityListFil 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 + `)` @@ -356,6 +364,7 @@ func (db *DB) DeleteVulnerabilitiesByFilterForAccess(filter VulnerabilityListFil if err := tx.Commit(); err != nil { return 0, fmt.Errorf("提交事务失败: %w", err) } + db.refreshAssetRiskCacheForConversationsBestEffort(affectedConversations...) return deleted, nil } @@ -366,6 +375,8 @@ func (db *DB) DeleteVulnerability(id string) error { 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 { @@ -377,9 +388,29 @@ func (db *DB) DeleteVulnerability(id string) error { 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{})