1use crate::diff::result::VulnerabilityDetail;
7use serde::{Deserialize, Serialize};
8use std::collections::HashMap;
9
10#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
12pub enum VulnGroupStatus {
13 Introduced,
15 Resolved,
17 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#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct VulnerabilityGroup {
34 pub component_id: String,
36 pub component_name: String,
38 pub component_version: Option<String>,
40 pub vulnerabilities: Vec<VulnerabilityDetail>,
42 pub max_severity: String,
44 pub max_cvss: Option<f32>,
46 pub severity_counts: HashMap<String, usize>,
48 pub status: VulnGroupStatus,
50 pub has_kev: bool,
52 pub has_ransomware_kev: bool,
54}
55
56impl VulnerabilityGroup {
57 #[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 pub fn add_vulnerability(&mut self, vuln: VulnerabilityDetail) {
76 *self
78 .severity_counts
79 .entry(vuln.severity.clone())
80 .or_insert(0) += 1;
81
82 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 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 if self.component_version.is_none() {
96 self.component_version.clone_from(&vuln.version);
97 }
98
99 if vuln.is_kev {
101 self.has_kev = true;
102 }
103
104 if vuln.is_ransomware {
106 self.has_ransomware_kev = true;
107 }
108
109 self.vulnerabilities.push(vuln);
110 }
111
112 #[must_use]
114 pub fn vuln_count(&self) -> usize {
115 self.vulnerabilities.len()
116 }
117
118 #[must_use]
120 pub fn has_critical(&self) -> bool {
121 self.severity_counts.get("Critical").copied().unwrap_or(0) > 0
122 }
123
124 #[must_use]
126 pub fn has_high(&self) -> bool {
127 self.severity_counts.get("High").copied().unwrap_or(0) > 0
128 }
129
130 #[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
162fn 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#[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 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#[derive(Debug, Clone, Default, Serialize, Deserialize)]
211pub struct VulnerabilityGroupedView {
212 pub introduced_groups: Vec<VulnerabilityGroup>,
214 pub resolved_groups: Vec<VulnerabilityGroup>,
216 pub persistent_groups: Vec<VulnerabilityGroup>,
218}
219
220impl VulnerabilityGroupedView {
221 #[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 #[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 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 #[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 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 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}