Skip to main content

sbom_tools/diff/changes/
vuln_grouping.rs

1//! Vulnerability grouping by root cause component.
2//!
3//! This module provides functionality to group vulnerabilities by the component
4//! that introduces them, reducing noise and showing the true scope of security issues.
5
6use crate::diff::result::VulnerabilityDetail;
7use serde::{Deserialize, Serialize};
8use std::collections::HashMap;
9
10/// Status of a vulnerability group
11#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
12pub enum VulnGroupStatus {
13    /// Newly introduced vulnerabilities
14    Introduced,
15    /// Resolved vulnerabilities
16    Resolved,
17    /// Persistent vulnerabilities (present in both old and new)
18    Persistent,
19}
20
21impl std::fmt::Display for VulnGroupStatus {
22    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
23        match self {
24            Self::Introduced => write!(f, "Introduced"),
25            Self::Resolved => write!(f, "Resolved"),
26            Self::Persistent => write!(f, "Persistent"),
27        }
28    }
29}
30
31/// A group of vulnerabilities sharing the same root cause component
32#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct VulnerabilityGroup {
34    /// Root cause component ID
35    pub component_id: String,
36    /// Component name
37    pub component_name: String,
38    /// Component version (if available)
39    pub component_version: Option<String>,
40    /// Vulnerabilities in this group
41    pub vulnerabilities: Vec<VulnerabilityDetail>,
42    /// Maximum severity in the group
43    pub max_severity: String,
44    /// Maximum CVSS score in the group
45    pub max_cvss: Option<f32>,
46    /// Count by severity level
47    pub severity_counts: HashMap<String, usize>,
48    /// Group status (Introduced, Resolved, Persistent)
49    pub status: VulnGroupStatus,
50    /// Whether any vulnerability is in KEV catalog
51    pub has_kev: bool,
52    /// Whether any vulnerability is ransomware-related
53    pub has_ransomware_kev: bool,
54}
55
56impl VulnerabilityGroup {
57    /// Create a new empty group for a component
58    #[must_use]
59    pub fn new(component_id: String, component_name: String, status: VulnGroupStatus) -> Self {
60        Self {
61            component_id,
62            component_name,
63            component_version: None,
64            vulnerabilities: Vec::new(),
65            max_severity: "Unknown".to_string(),
66            max_cvss: None,
67            severity_counts: HashMap::new(),
68            status,
69            has_kev: false,
70            has_ransomware_kev: false,
71        }
72    }
73
74    /// Add a vulnerability to the group
75    pub fn add_vulnerability(&mut self, vuln: VulnerabilityDetail) {
76        // Update severity counts
77        *self
78            .severity_counts
79            .entry(vuln.severity.clone())
80            .or_insert(0) += 1;
81
82        // Update max severity (priority: Critical > High > Medium > Low > Unknown)
83        let vuln_priority = severity_priority(&vuln.severity);
84        let current_priority = severity_priority(&self.max_severity);
85        if vuln_priority < current_priority {
86            self.max_severity.clone_from(&vuln.severity);
87        }
88
89        // Update max CVSS
90        if let Some(score) = vuln.cvss_score {
91            self.max_cvss = Some(self.max_cvss.map_or(score, |c| c.max(score)));
92        }
93
94        // Update version from first vulnerability with version
95        if self.component_version.is_none() {
96            self.component_version.clone_from(&vuln.version);
97        }
98
99        // Propagate the KEV (actively exploited) flag to the group.
100        if vuln.is_kev {
101            self.has_kev = true;
102        }
103
104        // Propagate the ransomware-campaign flag to the group.
105        if vuln.is_ransomware {
106            self.has_ransomware_kev = true;
107        }
108
109        self.vulnerabilities.push(vuln);
110    }
111
112    /// Get total vulnerability count
113    #[must_use]
114    pub fn vuln_count(&self) -> usize {
115        self.vulnerabilities.len()
116    }
117
118    /// Check if group has any critical vulnerabilities
119    #[must_use]
120    pub fn has_critical(&self) -> bool {
121        self.severity_counts.get("Critical").copied().unwrap_or(0) > 0
122    }
123
124    /// Check if group has any high severity vulnerabilities
125    #[must_use]
126    pub fn has_high(&self) -> bool {
127        self.severity_counts.get("High").copied().unwrap_or(0) > 0
128    }
129
130    /// Get summary line for display
131    #[must_use]
132    pub fn summary_line(&self) -> String {
133        let version_str = self
134            .component_version
135            .as_ref()
136            .map(|v| format!("@{v}"))
137            .unwrap_or_default();
138
139        let severity_badges: Vec<String> = ["Critical", "High", "Medium", "Low"]
140            .iter()
141            .filter_map(|sev| {
142                self.severity_counts.get(*sev).and_then(|&count| {
143                    if count > 0 {
144                        Some(format!("{}:{}", &sev[..1], count))
145                    } else {
146                        None
147                    }
148                })
149            })
150            .collect();
151
152        format!(
153            "{}{}: {} CVEs [{}]",
154            self.component_name,
155            version_str,
156            self.vuln_count(),
157            severity_badges.join(" ")
158        )
159    }
160}
161
162/// Get priority value for severity (lower = more severe)
163fn severity_priority(severity: &str) -> u8 {
164    match severity.to_lowercase().as_str() {
165        "critical" => 0,
166        "high" => 1,
167        "medium" => 2,
168        "low" => 3,
169        "info" => 4,
170        "none" => 5,
171        _ => 6,
172    }
173}
174
175/// Group vulnerabilities by component
176#[must_use]
177pub fn group_vulnerabilities(
178    vulns: &[VulnerabilityDetail],
179    status: VulnGroupStatus,
180) -> Vec<VulnerabilityGroup> {
181    let mut groups: HashMap<String, VulnerabilityGroup> = HashMap::new();
182
183    for vuln in vulns {
184        let group = groups.entry(vuln.component_id.clone()).or_insert_with(|| {
185            VulnerabilityGroup::new(
186                vuln.component_id.clone(),
187                vuln.component_name.clone(),
188                status,
189            )
190        });
191
192        group.add_vulnerability(vuln.clone());
193    }
194
195    // Sort groups by severity (most severe first), then by count
196    let mut result: Vec<_> = groups.into_values().collect();
197    result.sort_by(|a, b| {
198        let sev_cmp = severity_priority(&a.max_severity).cmp(&severity_priority(&b.max_severity));
199        if sev_cmp == std::cmp::Ordering::Equal {
200            b.vuln_count().cmp(&a.vuln_count())
201        } else {
202            sev_cmp
203        }
204    });
205
206    result
207}
208
209/// Grouped view of vulnerability changes
210#[derive(Debug, Clone, Default, Serialize, Deserialize)]
211pub struct VulnerabilityGroupedView {
212    /// Groups of introduced vulnerabilities
213    pub introduced_groups: Vec<VulnerabilityGroup>,
214    /// Groups of resolved vulnerabilities
215    pub resolved_groups: Vec<VulnerabilityGroup>,
216    /// Groups of persistent vulnerabilities
217    pub persistent_groups: Vec<VulnerabilityGroup>,
218}
219
220impl VulnerabilityGroupedView {
221    /// Create grouped view from vulnerability lists
222    #[must_use]
223    pub fn from_changes(
224        introduced: &[VulnerabilityDetail],
225        resolved: &[VulnerabilityDetail],
226        persistent: &[VulnerabilityDetail],
227    ) -> Self {
228        Self {
229            introduced_groups: group_vulnerabilities(introduced, VulnGroupStatus::Introduced),
230            resolved_groups: group_vulnerabilities(resolved, VulnGroupStatus::Resolved),
231            persistent_groups: group_vulnerabilities(persistent, VulnGroupStatus::Persistent),
232        }
233    }
234
235    /// Get total group count
236    #[must_use]
237    pub fn total_groups(&self) -> usize {
238        self.introduced_groups.len() + self.resolved_groups.len() + self.persistent_groups.len()
239    }
240
241    /// Get total vulnerability count across all groups
242    pub fn total_vulns(&self) -> usize {
243        self.introduced_groups
244            .iter()
245            .map(VulnerabilityGroup::vuln_count)
246            .sum::<usize>()
247            + self
248                .resolved_groups
249                .iter()
250                .map(VulnerabilityGroup::vuln_count)
251                .sum::<usize>()
252            + self
253                .persistent_groups
254                .iter()
255                .map(VulnerabilityGroup::vuln_count)
256                .sum::<usize>()
257    }
258
259    /// Check if any group has KEV vulnerabilities
260    #[must_use]
261    pub fn has_any_kev(&self) -> bool {
262        self.introduced_groups.iter().any(|g| g.has_kev)
263            || self.resolved_groups.iter().any(|g| g.has_kev)
264            || self.persistent_groups.iter().any(|g| g.has_kev)
265    }
266}
267
268#[cfg(test)]
269mod tests {
270    use super::*;
271
272    fn make_vuln(id: &str, component_id: &str, severity: &str) -> VulnerabilityDetail {
273        VulnerabilityDetail {
274            id: id.to_string(),
275            source: "OSV".to_string(),
276            severity: severity.to_string(),
277            cvss_score: None,
278            component_id: component_id.to_string(),
279            component_canonical_id: None,
280            component_ref: None,
281            component_name: format!("{}-pkg", component_id),
282            version: Some("1.0.0".to_string()),
283            description: None,
284            remediation: None,
285            is_kev: false,
286            is_ransomware: false,
287            epss_score: None,
288            cwes: Vec::new(),
289            component_depth: None,
290            published_date: None,
291            kev_due_date: None,
292            days_since_published: None,
293            days_until_due: None,
294            vex_state: None,
295            vex_justification: None,
296            vex_impact_statement: None,
297        }
298    }
299
300    #[test]
301    fn test_group_vulnerabilities() {
302        let vulns = vec![
303            make_vuln("CVE-2024-0001", "lodash", "Critical"),
304            make_vuln("CVE-2024-0002", "lodash", "High"),
305            make_vuln("CVE-2024-0003", "lodash", "High"),
306            make_vuln("CVE-2024-0004", "express", "Medium"),
307        ];
308
309        let groups = group_vulnerabilities(&vulns, VulnGroupStatus::Introduced);
310
311        assert_eq!(groups.len(), 2);
312
313        // lodash should be first (Critical severity)
314        assert_eq!(groups[0].component_id, "lodash");
315        assert_eq!(groups[0].vuln_count(), 3);
316        assert_eq!(groups[0].max_severity, "Critical");
317        assert_eq!(groups[0].severity_counts.get("Critical"), Some(&1));
318        assert_eq!(groups[0].severity_counts.get("High"), Some(&2));
319
320        // express should be second
321        assert_eq!(groups[1].component_id, "express");
322        assert_eq!(groups[1].vuln_count(), 1);
323    }
324
325    #[test]
326    fn test_group_propagates_kev_flag() {
327        let mut kev_vuln = make_vuln("CVE-2021-44228", "log4j", "Critical");
328        kev_vuln.is_kev = true;
329        let vulns = vec![kev_vuln, make_vuln("CVE-2024-0009", "lodash", "High")];
330
331        let groups = group_vulnerabilities(&vulns, VulnGroupStatus::Introduced);
332
333        let log4j = groups
334            .iter()
335            .find(|g| g.component_id == "log4j")
336            .expect("log4j group present");
337        assert!(log4j.has_kev, "group with a KEV vuln must report has_kev");
338
339        let lodash = groups
340            .iter()
341            .find(|g| g.component_id == "lodash")
342            .expect("lodash group present");
343        assert!(
344            !lodash.has_kev,
345            "group without KEV vulns must not report has_kev"
346        );
347    }
348
349    #[test]
350    fn test_group_propagates_ransomware_flag() {
351        let mut ransomware_vuln = make_vuln("CVE-2021-44228", "log4j", "Critical");
352        ransomware_vuln.is_kev = true;
353        ransomware_vuln.is_ransomware = true;
354        let vulns = vec![
355            ransomware_vuln,
356            make_vuln("CVE-2024-0009", "lodash", "High"),
357        ];
358
359        let groups = group_vulnerabilities(&vulns, VulnGroupStatus::Introduced);
360
361        let log4j = groups
362            .iter()
363            .find(|g| g.component_id == "log4j")
364            .expect("log4j group present");
365        assert!(
366            log4j.has_ransomware_kev,
367            "group with a ransomware-KEV vuln must report has_ransomware_kev"
368        );
369
370        let lodash = groups
371            .iter()
372            .find(|g| g.component_id == "lodash")
373            .expect("lodash group present");
374        assert!(
375            !lodash.has_ransomware_kev,
376            "group without ransomware vulns must not report has_ransomware_kev"
377        );
378    }
379
380    #[test]
381    fn test_grouped_view() {
382        let introduced = vec![
383            make_vuln("CVE-2024-0001", "lodash", "High"),
384            make_vuln("CVE-2024-0002", "lodash", "Medium"),
385        ];
386        let resolved = vec![make_vuln("CVE-2024-0003", "old-dep", "Critical")];
387        let persistent = vec![];
388
389        let view = VulnerabilityGroupedView::from_changes(&introduced, &resolved, &persistent);
390
391        assert_eq!(view.total_groups(), 2);
392        assert_eq!(view.total_vulns(), 3);
393        assert_eq!(view.introduced_groups.len(), 1);
394        assert_eq!(view.resolved_groups.len(), 1);
395    }
396
397    #[test]
398    fn test_summary_line() {
399        let mut group = VulnerabilityGroup::new(
400            "lodash".to_string(),
401            "lodash".to_string(),
402            VulnGroupStatus::Introduced,
403        );
404        group.add_vulnerability(make_vuln("CVE-1", "lodash", "Critical"));
405        group.add_vulnerability(make_vuln("CVE-2", "lodash", "High"));
406        group.add_vulnerability(make_vuln("CVE-3", "lodash", "High"));
407
408        let summary = group.summary_line();
409        assert!(summary.contains("lodash"));
410        assert!(summary.contains("3 CVEs"));
411        assert!(summary.contains("C:1"));
412        assert!(summary.contains("H:2"));
413    }
414}