mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-08-01 08:37:41 +02:00
Add files via upload
This commit is contained in:
+150
-8
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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" {
|
||||
|
||||
@@ -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{})
|
||||
|
||||
Reference in New Issue
Block a user