diff options
| author | Adam Janovsky | 2022-02-16 19:46:15 +0100 |
|---|---|---|
| committer | Adam Janovsky | 2022-02-16 19:46:15 +0100 |
| commit | 52f04fbfe926701b35de67d02fbc3d03c57138bc (patch) | |
| tree | abe95649a5eb4a9387be9b9349890b540a85383b | |
| parent | 4cd539293c85958c6fe115d3909d8ffc1c0b7af7 (diff) | |
| download | sec-certs-52f04fbfe926701b35de67d02fbc3d03c57138bc.tar.gz sec-certs-52f04fbfe926701b35de67d02fbc3d03c57138bc.tar.zst sec-certs-52f04fbfe926701b35de67d02fbc3d03c57138bc.zip | |
cpe matches list->set, heuristic type hints
| -rw-r--r-- | sec_certs/dataset/dataset.py | 7 | ||||
| -rw-r--r-- | sec_certs/dataset/fips.py | 2 | ||||
| -rw-r--r-- | sec_certs/helpers.py | 7 | ||||
| -rw-r--r-- | sec_certs/model/cpe_matching.py | 46 | ||||
| -rw-r--r-- | sec_certs/sample/certificate.py | 14 | ||||
| -rw-r--r-- | sec_certs/sample/common_criteria.py | 42 | ||||
| -rw-r--r-- | sec_certs/sample/fips.py | 29 | ||||
| -rw-r--r-- | tests/test_cc_heuristics.py | 4 |
8 files changed, 71 insertions, 80 deletions
diff --git a/sec_certs/dataset/dataset.py b/sec_certs/dataset/dataset.py index 2789d5e9..56b559e1 100644 --- a/sec_certs/dataset/dataset.py +++ b/sec_certs/dataset/dataset.py @@ -191,11 +191,6 @@ class Dataset(Generic[CertSubType], ABC): cve_dataset.build_lookup_dict(use_nist_cpe_matching_dict, self.nist_cve_cpe_matching_dset_path) return cve_dataset - def _compute_candidate_versions(self) -> None: - logger.info("Computing heuristics: possible product versions in sample name") - for cert in cast(Iterator[Certificate], self): - cert.compute_heuristics_version() - def _compute_cpe_matches( self, download_fresh_cpes: bool = False ) -> Tuple[CPEClassifier, CPEDataset, Optional[CVEDataset]]: @@ -230,6 +225,7 @@ class Dataset(Generic[CertSubType], ABC): clf = CPEClassifier(config.cpe_matching_threshold, config.cpe_n_max_matches) clf.fit([x for x in cpe_dset if filter_condition(x)]) + cert: CertSubType for cert in helpers.tqdm(self, desc="Predicting CPE matches with the classifier"): cert.compute_heuristics_cpe_match(clf) @@ -237,7 +233,6 @@ class Dataset(Generic[CertSubType], ABC): @serialize def compute_cpe_heuristics(self) -> Tuple[CPEClassifier, CPEDataset, Optional[CVEDataset]]: - self._compute_candidate_versions() return self._compute_cpe_matches() def to_label_studio_json(self, output_path: Union[str, Path]) -> None: diff --git a/sec_certs/dataset/fips.py b/sec_certs/dataset/fips.py index 2148038f..91926211 100644 --- a/sec_certs/dataset/fips.py +++ b/sec_certs/dataset/fips.py @@ -304,7 +304,7 @@ class FIPSDataset(Dataset[FIPSCertificate], ComplexSerializableType): ) cert: FIPSCertificate for cert in self.certs.values(): - cert.heuristics = FIPSCertificate.FIPSHeuristics(None, [], [], 0) + cert.heuristics = FIPSCertificate.FIPSHeuristics(dict(), [], [], 0) self.match_algs() diff --git a/sec_certs/helpers.py b/sec_certs/helpers.py index 6c775e73..41ac39f1 100644 --- a/sec_certs/helpers.py +++ b/sec_certs/helpers.py @@ -34,6 +34,7 @@ from sec_certs.constants import ( logger = logging.getLogger(__name__) +# TODO: Once typehints in tqdm are implemented, we should use them: https://github.com/tqdm/tqdm/issues/260 def tqdm(*args, **kwargs): if "disable" in kwargs: return tqdm_original(*args, **kwargs) @@ -1035,7 +1036,7 @@ def gen_dict_extract(dct: Dict, searched_key: Hashable = "count") -> Generator[A yield key, result -def compute_heuristics_version(cert_name: str) -> List[str]: +def compute_heuristics_version(cert_name: str) -> Set[str]: """ Will extract possible versions from the name of sample """ @@ -1060,10 +1061,10 @@ def compute_heuristics_version(cert_name: str) -> List[str]: # return identified_versions if identified_versions else ['-'] if not matches: - return ["-"] + return {"-"} matched = [re.search(normalizer, x) for x in matches] - return [x.group() for x in matched if x is not None] + return {x.group() for x in matched if x is not None} def tokenize_dataset(dset: List[str], keywords: Set[str]) -> np.ndarray: diff --git a/sec_certs/model/cpe_matching.py b/sec_certs/model/cpe_matching.py index 79cba635..a261aedb 100644 --- a/sec_certs/model/cpe_matching.py +++ b/sec_certs/model/cpe_matching.py @@ -67,7 +67,7 @@ class CPEClassifier(BaseEstimator): else: self.vendor_version_to_cpe_[(cpe.vendor, cpe.version)].add(cpe) - def predict(self, X: List[Tuple[str, str, str]]) -> List[Optional[List[str]]]: + def predict(self, X: List[Tuple[str, str, str]]) -> List[Optional[Set[str]]]: """ Will predict CPE uris for List of Tuples (vendor, product name, identified versions in product name) @param X: tuples (vendor, product name, identified versions in product name) @@ -79,10 +79,10 @@ class CPEClassifier(BaseEstimator): self, vendor: Optional[str], product_name: str, - versions: List[str], + versions: Set[str], relax_version: bool = False, relax_title: bool = False, - ) -> Optional[List[str]]: + ) -> Optional[Set[str]]: """ Predict List of CPE uris for triplet (vendor, product_name, list_of_version). The prediction is made as follows: 1. Sanitize all strings @@ -109,9 +109,9 @@ class CPEClassifier(BaseEstimator): ] threshold = self.match_threshold if not relax_version else 100 final_matches_aux: List[Tuple[float, CPE]] = list(filter(lambda x: x[0] >= threshold, zip(ratings, candidates))) - final_matches: Optional[List[str]] = [ - x[1].uri for x in final_matches_aux[: self.n_max_matches] if x[1].uri is not None - ] + final_matches: Optional[Set[str]] = set( + [x[1].uri for x in final_matches_aux[: self.n_max_matches] if x[1].uri is not None] + ) if not relax_title and not final_matches: final_matches = self.predict_single_cert( @@ -120,7 +120,7 @@ class CPEClassifier(BaseEstimator): if not relax_version and not final_matches: final_matches = self.predict_single_cert( - vendor, product_name, ["-"], relax_version=True, relax_title=relax_title + vendor, product_name, {"-"}, relax_version=True, relax_title=relax_title ) return final_matches if final_matches else None @@ -129,8 +129,8 @@ class CPEClassifier(BaseEstimator): self, cpe: CPE, product_name: str, - candidate_vendors: Optional[List[str]], - versions: List[str], + candidate_vendors: Optional[Set[str]], + versions: Set[str], relax_title: bool = False, ) -> float: """ @@ -183,13 +183,13 @@ class CPEClassifier(BaseEstimator): return string.replace("®", "").replace("™", "") @staticmethod - def _strip_manufacturer_and_version(string: str, manufacturers: Optional[List[str]], versions: List[str]) -> str: - to_strip = versions + manufacturers if manufacturers else versions + def _strip_manufacturer_and_version(string: str, manufacturers: Optional[Set[str]], versions: Set[str]) -> str: + to_strip = versions | manufacturers if manufacturers else versions for x in to_strip: string = string.lower().replace(CPEClassifier._replace_special_chars_with_space(x.lower()), "").strip() return string - def _process_manufacturer(self, manufacturer: str, result: Set) -> Optional[List[str]]: + def _process_manufacturer(self, manufacturer: str, result: Set) -> Set[str]: tokenized = manufacturer.split() if tokenized[0] in self.vendors_: result.add(tokenized[0]) @@ -208,20 +208,20 @@ class CPEClassifier(BaseEstimator): result.add("athena-scs") if tokenized[0] == "the" and not result: candidate_result = self.get_candidate_list_of_vendors(" ".join(tokenized[1:])) - return list(candidate_result) if candidate_result else None + return set(candidate_result) if candidate_result else set() - return list(result) if result else None + return set(result) if result else set() - def get_candidate_list_of_vendors(self, manufacturer: Optional[str]) -> Optional[List[str]]: + def get_candidate_list_of_vendors(self, manufacturer: Optional[str]) -> Set[str]: """ Given manufacturer name, this method will find list of plausible vendors from CPE dataset that are likely related. @param manufacturer: manufacturer @return: List of related manufacturers, None if nothing relevant is found. """ + result: Set[str] = set() if not manufacturer: - return None + return result - result: Set = set() splits = re.compile(r"[,/]").findall(manufacturer) if splits: @@ -229,8 +229,8 @@ class CPEClassifier(BaseEstimator): itertools.chain.from_iterable([[x.strip() for x in manufacturer.split(s)] for s in splits]) ) result_aux = [self.get_candidate_list_of_vendors(x) for x in vendor_tokens] - result_used = list(set(itertools.chain.from_iterable([x for x in result_aux if x]))) - return result_used if result_used else None + result_used = set(set(itertools.chain.from_iterable([x for x in result_aux if x]))) + return result_used if result_used else set() if manufacturer in self.vendors_: result.add(manufacturer) @@ -238,7 +238,7 @@ class CPEClassifier(BaseEstimator): return self._process_manufacturer(manufacturer, result) def get_candidate_vendor_version_pairs( - self, cert_candidate_cpe_vendors: Optional[List[str]], cert_candidate_versions: List[str] + self, cert_candidate_cpe_vendors: Set[str], cert_candidate_versions: Set[str] ) -> Optional[List[Tuple[str, str]]]: """ Given parameters, will return Pairs (cpe_vendor, cpe_version) that are relevant to a given sample @@ -247,7 +247,7 @@ class CPEClassifier(BaseEstimator): @return: List of tuples (cpe_vendor, cpe_version) that can be used in the lookup table to search the CPE dataset. """ - def is_cpe_version_among_cert_versions(cpe_version: Optional[str], cert_versions: List[str]) -> bool: + def is_cpe_version_among_cert_versions(cpe_version: Optional[str], cert_versions: Set[str]) -> bool: def simple_startswith(seeked_version: str, checked_string: str) -> bool: if seeked_version == checked_string: return True @@ -278,9 +278,7 @@ class CPEClassifier(BaseEstimator): candidate_vendor_version_pairs.extend([(vendor, x) for x in matched_cpe_versions]) return candidate_vendor_version_pairs - def get_candidate_cpe_matches( - self, candidate_vendors: Optional[List[str]], candidate_versions: List[str] - ) -> List[CPE]: + def get_candidate_cpe_matches(self, candidate_vendors: Set[str], candidate_versions: Set[str]) -> List[CPE]: """ Given List of candidate vendors and candidate versions found in certificate, candidate CPE matches are found @param candidate_vendors: List of version strings that were found in the certificate diff --git a/sec_certs/sample/certificate.py b/sec_certs/sample/certificate.py index a7f34d2a..c310dae2 100644 --- a/sec_certs/sample/certificate.py +++ b/sec_certs/sample/certificate.py @@ -4,7 +4,7 @@ import json import logging from abc import ABC, abstractmethod from pathlib import Path -from typing import Any, Dict, Generic, Type, TypeVar, Union +from typing import Any, Dict, Generic, Optional, Set, Type, TypeVar, Union from sec_certs.dataset.cve import CVEDataset from sec_certs.model.cpe_matching import CPEClassifier @@ -13,10 +13,16 @@ from sec_certs.serialization.json import ComplexSerializableType, CustomJSONDeco logger = logging.getLogger(__name__) T = TypeVar("T", bound="Certificate") +H = TypeVar("H", bound="Heuristics") -class Certificate(Generic[T], ABC, ComplexSerializableType): - heuristics: Any +class Heuristics: + cpe_matches: Optional[Set[str]] + related_cves: Optional[Set[str]] + + +class Certificate(Generic[T, H], ABC, ComplexSerializableType): + heuristics: H def __init__(self, *args, **kwargs): pass @@ -59,7 +65,7 @@ class Certificate(Generic[T], ABC, ComplexSerializableType): return json.load(handle, cls=CustomJSONDecoder) @abstractmethod - def compute_heuristics_version(self) -> None: + def _compute_heuristics_version(self) -> None: raise NotImplementedError("Not meant to be implemented") @abstractmethod diff --git a/sec_certs/sample/common_criteria.py b/sec_certs/sample/common_criteria.py index 5dda7326..42468aab 100644 --- a/sec_certs/sample/common_criteria.py +++ b/sec_certs/sample/common_criteria.py @@ -12,7 +12,7 @@ from bs4 import Tag from sec_certs import constants as constants from sec_certs import helpers from sec_certs.model.cpe_matching import CPEClassifier -from sec_certs.sample.certificate import Certificate, logger +from sec_certs.sample.certificate import Certificate, Heuristics, logger from sec_certs.sample.protection_profile import ProtectionProfile from sec_certs.serialization.json import ComplexSerializableType from sec_certs.serialization.pandas import PandasSerializableType @@ -26,7 +26,11 @@ HEADERS = { } -class CommonCriteriaCert(Certificate["CommonCriteriaCert"], PandasSerializableType, ComplexSerializableType): +class CommonCriteriaCert( + Certificate["CommonCriteriaCert", "CommonCriteriaCert.CCHeuristics"], + PandasSerializableType, + ComplexSerializableType, +): cc_url = "http://www.commoncriteriaportal.org" empty_st_url = "http://www.commoncriteriaportal.org/files/epfiles/" @@ -198,8 +202,8 @@ class CommonCriteriaCert(Certificate["CommonCriteriaCert"], PandasSerializableTy return processed if (processed := self.processed_cert_id) else self.keywords_cert_id @dataclass - class CCHeuristics(ComplexSerializableType): - extracted_versions: Optional[List[str]] = field(default=None) + class CCHeuristics(Heuristics, ComplexSerializableType): + extracted_versions: Optional[Set[str]] = field(default=None) cpe_matches: Optional[Set[str]] = field(default=None) verified_cpe_matches: Optional[Set[str]] = field(default=None) related_cves: Optional[Set[str]] = field(default=None) @@ -209,16 +213,10 @@ class CommonCriteriaCert(Certificate["CommonCriteriaCert"], PandasSerializableTy indirectly_affected_by: Optional[Set[str]] = field(default=None) directly_affecting: Optional[Set[str]] = field(default=None) indirectly_affecting: Optional[Set[str]] = field(default=None) - cpe_candidate_vendors: Optional[List[str]] = field(init=False) @property def serialized_attributes(self) -> List[str]: - all_vars = copy.deepcopy(super().serialized_attributes) - all_vars.remove("cpe_candidate_vendors") - return all_vars - - def __post_init__(self) -> None: - self.cpe_candidate_vendors = None + return copy.deepcopy(super().serialized_attributes) pandas_columns: ClassVar[List[str]] = [ "dgst", @@ -284,18 +282,9 @@ class CommonCriteriaCert(Certificate["CommonCriteriaCert"], PandasSerializableTy self.manufacturer_web = helpers.sanitize_link(manufacturer_web) self.protection_profiles = protection_profiles self.maintenance_updates = maintenance_updates - - if state is None: - state = self.InternalState() - self.state = state - - if pdf_data is None: - pdf_data = self.PdfData() - self.pdf_data = pdf_data - - if heuristics is None: - heuristics = self.CCHeuristics() - self.heuristics = heuristics + self.state = self.InternalState() if not state else state + self.pdf_data = self.PdfData() if not pdf_data else pdf_data + self.heuristics: "CommonCriteriaCert.CCHeuristics" = self.CCHeuristics() if not heuristics else heuristics @property def dgst(self) -> str: @@ -643,10 +632,13 @@ class CommonCriteriaCert(Certificate["CommonCriteriaCert"], PandasSerializableTy cert.state.errors.append(response) return cert - def compute_heuristics_version(self) -> None: + def _compute_heuristics_version(self) -> None: self.heuristics.extracted_versions = helpers.compute_heuristics_version(self.name) def compute_heuristics_cpe_match(self, cpe_classifier: CPEClassifier) -> None: + self._compute_heuristics_version() + assert self.heuristics.extracted_versions is not None + self.heuristics.cpe_matches = cpe_classifier.predict_single_cert( self.manufacturer, self.name, self.heuristics.extracted_versions ) @@ -770,6 +762,8 @@ class CommonCriteriaCert(Certificate["CommonCriteriaCert"], PandasSerializableTy return new_cert_id def get_cert_laboratory(self) -> str: + if not self.heuristics.cert_id: + raise ValueError("Cert ID was None but cert laboratory was to be computed based on its value.") cert_id = self.heuristics.cert_id.strip() if CommonCriteriaCert._is_anssi_cert(cert_id): diff --git a/sec_certs/sample/fips.py b/sec_certs/sample/fips.py index 562a4a3e..97fd30bd 100644 --- a/sec_certs/sample/fips.py +++ b/sec_certs/sample/fips.py @@ -17,12 +17,12 @@ from sec_certs.config.configuration import config from sec_certs.constants import LINE_SEPARATOR from sec_certs.helpers import fips_dgst, load_cert_file, normalize_match_string, save_modified_cert_file from sec_certs.model.cpe_matching import CPEClassifier -from sec_certs.sample.certificate import Certificate, logger +from sec_certs.sample.certificate import Certificate, Heuristics, logger from sec_certs.sample.cpe import CPE from sec_certs.serialization.json import ComplexSerializableType -class FIPSCertificate(Certificate["FIPSCertificate"], ComplexSerializableType): +class FIPSCertificate(Certificate["FIPSCertificate", "FIPSCertificate.FIPSHeuristics"], ComplexSerializableType): FIPS_BASE_URL: ClassVar[str] = "https://csrc.nist.gov" FIPS_MODULE_URL: ClassVar[ str @@ -167,17 +167,16 @@ class FIPSCertificate(Certificate["FIPSCertificate"], ComplexSerializableType): return str(self.cert_id) @dataclass(eq=True) - class FIPSHeuristics(ComplexSerializableType): - keywords: Optional[Dict[str, Dict]] + class FIPSHeuristics(Heuristics, ComplexSerializableType): + keywords: Dict[str, Dict] algorithms: List[Dict[str, Dict]] connections: List[str] unmatched_algs: int - extracted_versions: Optional[List[str]] = field(default=None) + extracted_versions: Optional[Set[str]] = field(default=None) cpe_matches: Optional[Set[str]] = field(default=None) verified_cpe_matches: Optional[Set[CPE]] = field(default=None) - related_cves: Optional[List[str]] = field(default=None) - cpe_candidate_vendors: Optional[List[str]] = field(init=False) + related_cves: Optional[Set[str]] = field(default=None) directly_affected_by: Optional[Set] = field(default=None) indirectly_affected_by: Optional[Set] = field(default=None) @@ -186,12 +185,7 @@ class FIPSCertificate(Certificate["FIPSCertificate"], ComplexSerializableType): @property def serialized_attributes(self) -> List[str]: - all_vars = copy.deepcopy(super().serialized_attributes) - all_vars.remove("cpe_candidate_vendors") - return all_vars - - def __post_init__(self) -> None: - self.cpe_candidate_vendors = None + return copy.deepcopy(super().serialized_attributes) @property def dgst(self) -> str: @@ -240,7 +234,7 @@ class FIPSCertificate(Certificate["FIPSCertificate"], ComplexSerializableType): self.cert_id = cert_id self.web_scan = web_scan self.pdf_scan = pdf_scan - self.heuristics = heuristics + self.heuristics: "FIPSCertificate.FIPSHeuristics" = heuristics self.state = state @classmethod @@ -544,7 +538,7 @@ class FIPSCertificate(Certificate["FIPSCertificate"], ComplexSerializableType): [] if not initialized else initialized.pdf_scan.algorithms, [], # connections ), - FIPSCertificate.FIPSHeuristics(None, [], [], 0), + FIPSCertificate.FIPSHeuristics(dict(), [], [], 0), state, ) @@ -840,7 +834,7 @@ class FIPSCertificate(Certificate["FIPSCertificate"], ComplexSerializableType): ) return vendor_split[0][:4] if len(vendor_split) > 0 else vendor - def compute_heuristics_version(self) -> None: + def _compute_heuristics_version(self) -> None: versions_for_extraction = "" if self.web_scan.module_name: versions_for_extraction += f" {self.web_scan.module_name}" @@ -851,6 +845,9 @@ class FIPSCertificate(Certificate["FIPSCertificate"], ComplexSerializableType): self.heuristics.extracted_versions = helpers.compute_heuristics_version(versions_for_extraction) def compute_heuristics_cpe_match(self, cpe_classifier: CPEClassifier) -> None: + self._compute_heuristics_version() + assert self.heuristics.extracted_versions is not None + if not self.web_scan.module_name: self.heuristics.cpe_matches = None else: diff --git a/tests/test_cc_heuristics.py b/tests/test_cc_heuristics.py index ec8a49ed..51ddad85 100644 --- a/tests/test_cc_heuristics.py +++ b/tests/test_cc_heuristics.py @@ -162,7 +162,7 @@ class TestCommonCriteriaHeuristics(TestCase): def test_version_extraction(self): self.assertEqual( self.cc_dset["ebd276cca70fd723"].heuristics.extracted_versions, - ["8.2"], + {"8.2"}, "The version extracted from the sample does not match the template", ) new_cert = CommonCriteriaCert( @@ -184,7 +184,7 @@ class TestCommonCriteriaHeuristics(TestCase): None, None, ) - new_cert.compute_heuristics_version() + new_cert._compute_heuristics_version() self.assertEqual( set(new_cert.heuristics.extracted_versions), {"5.4", "1.0"}, |
