mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-08-02 17:08:49 +02:00
Add files via upload
This commit is contained in:
+150
-8
@@ -13,6 +13,7 @@ import (
|
|||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
|
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
|
"go.uber.org/zap"
|
||||||
"golang.org/x/net/idna"
|
"golang.org/x/net/idna"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -50,6 +51,7 @@ type Asset struct {
|
|||||||
LastScanTaskID string `json:"last_scan_task_id,omitempty"`
|
LastScanTaskID string `json:"last_scan_task_id,omitempty"`
|
||||||
VulnerabilityCount int `json:"vulnerability_count"`
|
VulnerabilityCount int `json:"vulnerability_count"`
|
||||||
RiskLevel string `json:"risk_level"`
|
RiskLevel string `json:"risk_level"`
|
||||||
|
RiskScore int `json:"-"`
|
||||||
OwnerUserID string `json:"-"`
|
OwnerUserID string `json:"-"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -443,15 +445,15 @@ func assetWhere(filter AssetListFilter, access RBACListAccess) (string, []interf
|
|||||||
args = append(args, *filter.Port)
|
args = append(args, *filter.Port)
|
||||||
}
|
}
|
||||||
if filter.RiskLevel != "" {
|
if filter.RiskLevel != "" {
|
||||||
query += " AND " + assetRiskLevelExpr + " = ?"
|
query += " AND " + assetRiskLevelCachedExpr + " = ?"
|
||||||
args = append(args, strings.ToLower(strings.TrimSpace(filter.RiskLevel)))
|
args = append(args, strings.ToLower(strings.TrimSpace(filter.RiskLevel)))
|
||||||
}
|
}
|
||||||
if filter.MinVulnerabilities != nil {
|
if filter.MinVulnerabilities != nil {
|
||||||
query += " AND " + assetVulnerabilityCountExpr + " >= ?"
|
query += " AND " + assetVulnerabilityCountCachedExpr + " >= ?"
|
||||||
args = append(args, *filter.MinVulnerabilities)
|
args = append(args, *filter.MinVulnerabilities)
|
||||||
}
|
}
|
||||||
if filter.MaxVulnerabilities != nil {
|
if filter.MaxVulnerabilities != nil {
|
||||||
query += " AND " + assetVulnerabilityCountExpr + " <= ?"
|
query += " AND " + assetVulnerabilityCountCachedExpr + " <= ?"
|
||||||
args = append(args, *filter.MaxVulnerabilities)
|
args = append(args, *filter.MaxVulnerabilities)
|
||||||
}
|
}
|
||||||
for _, item := range []struct {
|
for _, item := range []struct {
|
||||||
@@ -573,20 +575,24 @@ const assetVulnerabilityMatchExpr = `(
|
|||||||
|
|
||||||
const assetVulnerabilityCountExpr = `(SELECT COUNT(DISTINCT v.id) FROM vulnerabilities v WHERE ` + 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)
|
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
|
FROM vulnerabilities v
|
||||||
WHERE LOWER(COALESCE(v.status,'open')) NOT IN ('fixed','false_positive','ignored') AND ` + assetVulnerabilityMatchExpr + `
|
WHERE LOWER(COALESCE(v.status,'open')) NOT IN ('fixed','false_positive','ignored') AND ` + assetVulnerabilityMatchExpr + `
|
||||||
),0)`
|
),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)`
|
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,
|
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.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,
|
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,''),
|
` + 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.
|
// 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.
|
// 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 {
|
if n == 0 {
|
||||||
return sql.ErrNoRows
|
return sql.ErrNoRows
|
||||||
}
|
}
|
||||||
|
if err := db.RefreshAssetRiskCache(id); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -629,6 +638,9 @@ func (db *DB) CompleteAssetScan(id, conversationID string, access RBACListAccess
|
|||||||
if n == 0 {
|
if n == 0 {
|
||||||
return sql.ErrNoRows
|
return sql.ErrNoRows
|
||||||
}
|
}
|
||||||
|
if err := db.RefreshAssetRiskCache(id); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -638,6 +650,136 @@ func (db *DB) BatchTaskBelongsToQueue(taskID, queueID string) bool {
|
|||||||
return err == nil && count > 0
|
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) {
|
func (db *DB) ListAssets(limit, offset int, filter AssetListFilter, access RBACListAccess) ([]*Asset, int, error) {
|
||||||
if limit < 1 {
|
if limit < 1 {
|
||||||
limit = 20
|
limit = 20
|
||||||
@@ -726,9 +868,9 @@ func assetOrderBy(sortBy, sortOrder string) string {
|
|||||||
case "port":
|
case "port":
|
||||||
expression = "assets.port"
|
expression = "assets.port"
|
||||||
case "vulnerability_count":
|
case "vulnerability_count":
|
||||||
expression = assetVulnerabilityCountExpr
|
expression = assetVulnerabilityCountCachedExpr
|
||||||
case "risk_level":
|
case "risk_level":
|
||||||
expression = assetRiskScoreExpr
|
expression = assetRiskScoreCachedExpr
|
||||||
default:
|
default:
|
||||||
expression = "assets.last_seen_at"
|
expression = "assets.last_seen_at"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -383,7 +383,12 @@ func TestAssetScanLinkReturnsTimeAndRelatedVulnerabilities(t *testing.T) {
|
|||||||
if linked.LastScanAt == nil || linked.LastScanConversationID != conv.ID || linked.VulnerabilityCount != 1 || linked.RiskLevel != "high" {
|
if linked.LastScanAt == nil || linked.LastScanConversationID != conv.ID || linked.VulnerabilityCount != 1 || linked.RiskLevel != "high" {
|
||||||
t.Fatalf("unexpected scan metadata: %#v", linked)
|
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)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
resolved, err := db.GetAsset(assets[0].ID, RBACListAccess{Scope: RBACScopeAll})
|
resolved, err := db.GetAsset(assets[0].ID, RBACListAccess{Scope: RBACScopeAll})
|
||||||
|
|||||||
@@ -1424,6 +1424,7 @@ type ProcessDetailsSummary struct {
|
|||||||
|
|
||||||
type ProcessDetailsToolExecution struct {
|
type ProcessDetailsToolExecution struct {
|
||||||
ProcessDetailID string `json:"processDetailId,omitempty"`
|
ProcessDetailID string `json:"processDetailId,omitempty"`
|
||||||
|
ResultDetailID string `json:"resultDetailId,omitempty"`
|
||||||
ToolName string `json:"toolName,omitempty"`
|
ToolName string `json:"toolName,omitempty"`
|
||||||
ToolCallID string `json:"toolCallId,omitempty"`
|
ToolCallID string `json:"toolCallId,omitempty"`
|
||||||
ExecutionID string `json:"executionId,omitempty"`
|
ExecutionID string `json:"executionId,omitempty"`
|
||||||
@@ -1552,6 +1553,7 @@ func (db *DB) GetProcessDetailsSummary(messageID string) (*ProcessDetailsSummary
|
|||||||
if toolCallID != "" {
|
if toolCallID != "" {
|
||||||
lastMatchedToolIndexByCallID[toolCallID] = idx
|
lastMatchedToolIndexByCallID[toolCallID] = idx
|
||||||
}
|
}
|
||||||
|
summary.ToolExecutions[idx].ResultDetailID = strings.TrimSpace(detailID)
|
||||||
if summary.ToolExecutions[idx].ToolName == "" {
|
if summary.ToolExecutions[idx].ToolName == "" {
|
||||||
summary.ToolExecutions[idx].ToolName = toolName
|
summary.ToolExecutions[idx].ToolName = toolName
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 '',
|
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 '',
|
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',
|
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,
|
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,
|
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
|
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_status ON assets(status);
|
||||||
CREATE INDEX IF NOT EXISTS idx_assets_owner ON assets(owner_user_id);
|
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_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_status ON projects(status);
|
||||||
CREATE INDEX IF NOT EXISTS idx_projects_updated_at ON projects(updated_at);
|
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_project_id ON project_facts(project_id);
|
||||||
@@ -977,7 +981,6 @@ func (db *DB) initTables() error {
|
|||||||
if _, err := db.Exec(createIndexes); err != nil {
|
if _, err := db.Exec(createIndexes); err != nil {
|
||||||
return fmt.Errorf("创建索引失败: %w", err)
|
return fmt.Errorf("创建索引失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
db.logger.Debug("数据库表初始化完成")
|
db.logger.Debug("数据库表初始化完成")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -1027,6 +1030,9 @@ func (db *DB) migrateAssetsTable() error {
|
|||||||
{"business_system", "ALTER TABLE assets ADD COLUMN business_system 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 ''"},
|
{"environment", "ALTER TABLE assets ADD COLUMN environment TEXT NOT NULL DEFAULT ''"},
|
||||||
{"criticality", "ALTER TABLE assets ADD COLUMN criticality 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 {
|
for _, column := range columns {
|
||||||
var count int
|
var count int
|
||||||
|
|||||||
@@ -22,10 +22,13 @@ func TestProcessDetailsSummaryDoesNotGuessIDLessResultsByOrder(t *testing.T) {
|
|||||||
{"toolName": "http-framework-test", "success": true},
|
{"toolName": "http-framework-test", "success": true},
|
||||||
{"toolName": "http-framework-test", "success": true},
|
{"toolName": "http-framework-test", "success": true},
|
||||||
}
|
}
|
||||||
|
var resultIDs []string
|
||||||
for _, result := range results {
|
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)
|
t.Fatalf("AddProcessDetail(tool_result): %v", err)
|
||||||
}
|
}
|
||||||
|
resultIDs = append(resultIDs, resultID)
|
||||||
}
|
}
|
||||||
|
|
||||||
summary, err := db.GetProcessDetailsSummary(messageID)
|
summary, err := db.GetProcessDetailsSummary(messageID)
|
||||||
@@ -39,6 +42,9 @@ func TestProcessDetailsSummaryDoesNotGuessIDLessResultsByOrder(t *testing.T) {
|
|||||||
if execution.Status != "completed" {
|
if execution.Status != "completed" {
|
||||||
t.Fatalf("execution %d status = %q, want completed", i, execution.Status)
|
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] {
|
for i, execution := range summary.ToolExecutions[2:4] {
|
||||||
if execution.Status != "result_missing" {
|
if execution.Status != "result_missing" {
|
||||||
|
|||||||
@@ -191,6 +191,7 @@ func (db *DB) CreateVulnerability(vuln *Vulnerability) (*Vulnerability, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("创建漏洞失败: %w", err)
|
return nil, fmt.Errorf("创建漏洞失败: %w", err)
|
||||||
}
|
}
|
||||||
|
db.refreshAssetRiskCacheForConversationsBestEffort(vuln.ConversationID)
|
||||||
return vuln, nil
|
return vuln, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -299,6 +300,8 @@ func (db *DB) CountVulnerabilitiesForAccess(filter VulnerabilityListFilter, acce
|
|||||||
// UpdateVulnerability 更新漏洞
|
// UpdateVulnerability 更新漏洞
|
||||||
func (db *DB) UpdateVulnerability(id string, vuln *Vulnerability) error {
|
func (db *DB) UpdateVulnerability(id string, vuln *Vulnerability) error {
|
||||||
vuln.UpdatedAt = time.Now()
|
vuln.UpdatedAt = time.Now()
|
||||||
|
var oldConversationID string
|
||||||
|
_ = db.QueryRow(`SELECT COALESCE(conversation_id,'') FROM vulnerabilities WHERE id = ?`, id).Scan(&oldConversationID)
|
||||||
|
|
||||||
query := `
|
query := `
|
||||||
UPDATE vulnerabilities
|
UPDATE vulnerabilities
|
||||||
@@ -318,6 +321,7 @@ func (db *DB) UpdateVulnerability(id string, vuln *Vulnerability) error {
|
|||||||
return fmt.Errorf("更新漏洞失败: %w", err)
|
return fmt.Errorf("更新漏洞失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
db.refreshAssetRiskCacheForConversationsBestEffort(oldConversationID, vuln.ConversationID)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -337,6 +341,10 @@ func (db *DB) DeleteVulnerabilitiesByFilterForAccess(filter VulnerabilityListFil
|
|||||||
args := []interface{}{}
|
args := []interface{}{}
|
||||||
where, args = filter.appendWhere(where, args)
|
where, args = filter.appendWhere(where, args)
|
||||||
where, args = appendVulnerabilityAccessFilter(where, args, access)
|
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
|
clearQuery := `UPDATE project_facts SET related_vulnerability_id = NULL
|
||||||
WHERE related_vulnerability_id IN (SELECT id FROM vulnerabilities ` + where + `)`
|
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 {
|
if err := tx.Commit(); err != nil {
|
||||||
return 0, fmt.Errorf("提交事务失败: %w", err)
|
return 0, fmt.Errorf("提交事务失败: %w", err)
|
||||||
}
|
}
|
||||||
|
db.refreshAssetRiskCacheForConversationsBestEffort(affectedConversations...)
|
||||||
return deleted, nil
|
return deleted, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -366,6 +375,8 @@ func (db *DB) DeleteVulnerability(id string) error {
|
|||||||
return fmt.Errorf("开启事务失败: %w", err)
|
return fmt.Errorf("开启事务失败: %w", err)
|
||||||
}
|
}
|
||||||
defer func() { _ = tx.Rollback() }()
|
defer func() { _ = tx.Rollback() }()
|
||||||
|
var conversationID string
|
||||||
|
_ = tx.QueryRow(`SELECT COALESCE(conversation_id,'') FROM vulnerabilities WHERE id = ?`, id).Scan(&conversationID)
|
||||||
|
|
||||||
// 删除漏洞前先解除项目事实中的关联,避免前端继续显示已删除漏洞的短 ID。
|
// 删除漏洞前先解除项目事实中的关联,避免前端继续显示已删除漏洞的短 ID。
|
||||||
if _, err := tx.Exec("UPDATE project_facts SET related_vulnerability_id = NULL WHERE related_vulnerability_id = ?", id); err != nil {
|
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 {
|
if err := tx.Commit(); err != nil {
|
||||||
return fmt.Errorf("提交事务失败: %w", err)
|
return fmt.Errorf("提交事务失败: %w", err)
|
||||||
}
|
}
|
||||||
|
db.refreshAssetRiskCacheForConversationsBestEffort(conversationID)
|
||||||
return nil
|
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 一致)
|
// GetVulnerabilityStats 获取漏洞统计(筛选条件与 ListVulnerabilities / CountVulnerabilities 一致)
|
||||||
func (db *DB) GetVulnerabilityStats(filter VulnerabilityListFilter) (map[string]interface{}, error) {
|
func (db *DB) GetVulnerabilityStats(filter VulnerabilityListFilter) (map[string]interface{}, error) {
|
||||||
return db.GetVulnerabilityStatsForAccess(filter, RBACListAccess{})
|
return db.GetVulnerabilityStatsForAccess(filter, RBACListAccess{})
|
||||||
|
|||||||
Reference in New Issue
Block a user