aboutsummaryrefslogtreecommitdiffhomepage
diff options
context:
space:
mode:
authorGeogeFI2022-04-09 13:56:08 +0200
committerGeogeFI2022-04-09 13:56:08 +0200
commite44481aaaaf3487a8bb6c3154ca71e892acf1386 (patch)
tree568fb909bebd0cbb4d743bba379cf97dcb86cf8c
parent8c76f33f43aaa89b28f83874d0f1eb96164a2dba (diff)
downloadsec-certs-e44481aaaaf3487a8bb6c3154ca71e892acf1386.tar.gz
sec-certs-e44481aaaaf3487a8bb6c3154ca71e892acf1386.tar.zst
sec-certs-e44481aaaaf3487a8bb6c3154ca71e892acf1386.zip
refactor: Renamed methods to suit sklearn philosophy
-rw-r--r--sec_certs/dataset/common_criteria.py6
-rw-r--r--sec_certs/model/dependency_vulnerability_finder.py22
2 files changed, 21 insertions, 7 deletions
diff --git a/sec_certs/dataset/common_criteria.py b/sec_certs/dataset/common_criteria.py
index 8e61d614..9e6d6027 100644
--- a/sec_certs/dataset/common_criteria.py
+++ b/sec_certs/dataset/common_criteria.py
@@ -686,11 +686,11 @@ class CCDataset(Dataset[CommonCriteriaCert], ComplexSerializableType):
cert.compute_heuristics_cert_id(self.all_cert_ids)
def _compute_dependency_vulnerabilities(self):
- cve_dependency_finder = DependencyVulnerabilityFinder(self.certs)
- cve_dependency_finder.fit()
+ cve_dependency_finder = DependencyVulnerabilityFinder()
+ cve_dependency_finder.fit(self.certs)
for dgst in self.certs:
- dependency_cve = cve_dependency_finder.get_dependency_vulnerabilities(dgst)
+ dependency_cve = cve_dependency_finder.predict_single_cert(dgst)
self.certs[dgst].heuristics.direct_dependency_cves = dependency_cve.direct_dependency_cves
self.certs[dgst].heuristics.indirect_dependency_cves = dependency_cve.indirect_dependency_cves
diff --git a/sec_certs/model/dependency_vulnerability_finder.py b/sec_certs/model/dependency_vulnerability_finder.py
index b2580891..2f66f30c 100644
--- a/sec_certs/model/dependency_vulnerability_finder.py
+++ b/sec_certs/model/dependency_vulnerability_finder.py
@@ -1,7 +1,7 @@
import logging
from dataclasses import dataclass, field
from enum import Enum
-from typing import Dict, Optional, Set
+from typing import Dict, List, Optional, Set
from sec_certs.sample.certificate import Certificate
from sec_certs.serialization.json import ComplexSerializableType
@@ -23,8 +23,12 @@ Vulnerabilities = Dict[str, Dict[str, Optional[Set[str]]]]
class DependencyVulnerabilityFinder:
- def __init__(self, certificates: Certificates):
+ def __init__(self):
self.vulnerabilities: Vulnerabilities = {}
+ self.certificates: Certificates = {}
+
+ def _overwrite_previous_state(self, certificates: Certificates) -> None:
+ self.vulnerabilities = {}
self.certificates = certificates
def _get_dataset_cert_ids_occurrences(self) -> Dict[str, int]:
@@ -72,7 +76,9 @@ class DependencyVulnerabilityFinder:
return vulnerabilities if vulnerabilities else None
- def fit(self) -> Vulnerabilities:
+ def fit(self, certificates: Certificates) -> Vulnerabilities:
+ self._overwrite_previous_state(certificates)
+
cert_id_occurrences = self._get_dataset_cert_ids_occurrences()
thrown_away_cert_counter = 0
@@ -99,7 +105,7 @@ class DependencyVulnerabilityFinder:
return self.vulnerabilities
- def get_dependency_vulnerabilities(self, dgst: str) -> DependencyCVE:
+ def predict_single_cert(self, dgst: str) -> DependencyCVE:
if not self.vulnerabilities.get(dgst):
return DependencyCVE(direct_dependency_cves=None, indirect_dependency_cves=None)
@@ -107,3 +113,11 @@ class DependencyVulnerabilityFinder:
self.vulnerabilities[dgst][DependencyType.DIRECT.value],
self.vulnerabilities[dgst][DependencyType.INDIRECT.value],
)
+
+ def predict(self, dgst_list: List[str]) -> Dict[str, DependencyCVE]:
+ cert_vulnerabilities = {}
+
+ for dgst in dgst_list:
+ cert_vulnerabilities[dgst] = self.predict_single_cert(dgst)
+
+ return cert_vulnerabilities