diff options
| author | GeogeFI | 2022-04-09 13:56:08 +0200 |
|---|---|---|
| committer | GeogeFI | 2022-04-09 13:56:08 +0200 |
| commit | e44481aaaaf3487a8bb6c3154ca71e892acf1386 (patch) | |
| tree | 568fb909bebd0cbb4d743bba379cf97dcb86cf8c | |
| parent | 8c76f33f43aaa89b28f83874d0f1eb96164a2dba (diff) | |
| download | sec-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.py | 6 | ||||
| -rw-r--r-- | sec_certs/model/dependency_vulnerability_finder.py | 22 |
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 |
