aboutsummaryrefslogtreecommitdiffhomepage
diff options
context:
space:
mode:
authorAdam Janovsky2022-02-16 19:46:15 +0100
committerAdam Janovsky2022-02-16 19:46:15 +0100
commit52f04fbfe926701b35de67d02fbc3d03c57138bc (patch)
treeabe95649a5eb4a9387be9b9349890b540a85383b
parent4cd539293c85958c6fe115d3909d8ffc1c0b7af7 (diff)
downloadsec-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.py7
-rw-r--r--sec_certs/dataset/fips.py2
-rw-r--r--sec_certs/helpers.py7
-rw-r--r--sec_certs/model/cpe_matching.py46
-rw-r--r--sec_certs/sample/certificate.py14
-rw-r--r--sec_certs/sample/common_criteria.py42
-rw-r--r--sec_certs/sample/fips.py29
-rw-r--r--tests/test_cc_heuristics.py4
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"},