aboutsummaryrefslogtreecommitdiffhomepage
path: root/sec_certs/dataset/dataset.py
diff options
context:
space:
mode:
Diffstat (limited to 'sec_certs/dataset/dataset.py')
-rw-r--r--sec_certs/dataset/dataset.py178
1 files changed, 102 insertions, 76 deletions
diff --git a/sec_certs/dataset/dataset.py b/sec_certs/dataset/dataset.py
index be06416c..eeae0601 100644
--- a/sec_certs/dataset/dataset.py
+++ b/sec_certs/dataset/dataset.py
@@ -1,36 +1,40 @@
-from datetime import datetime
-import logging
-from typing import Dict, Collection, Optional, Set, Union, List, Tuple, Mapping
-
+import itertools
import json
+import logging
from abc import ABC, abstractmethod
+from datetime import datetime
from pathlib import Path
-import itertools
+from typing import Collection, Dict, List, Mapping, Optional, Set, Tuple, Type, TypeVar, Union
import requests
-import sec_certs.helpers as helpers
import sec_certs.constants as constants
+import sec_certs.helpers as helpers
import sec_certs.parallel_processing as cert_processing
-from sec_certs.sample.cpe import CPE
-
-from sec_certs.sample.certificate import Certificate
-from sec_certs.serialization.json import ComplexSerializableType
from sec_certs.config.configuration import config
-from sec_certs.serialization.json import serialize
from sec_certs.dataset.cpe import CPEDataset
from sec_certs.dataset.cve import CVEDataset
from sec_certs.model.cpe_matching import CPEClassifier
+from sec_certs.sample.certificate import Certificate
+from sec_certs.sample.cpe import CPE
+from sec_certs.serialization.json import ComplexSerializableType, serialize
logger = logging.getLogger(__name__)
+T = TypeVar("T")
+
class Dataset(ABC):
- def __init__(self, certs: Mapping[str, 'Certificate'], root_dir: Path, name: str = 'dataset name',
- description: str = 'dataset_description'):
+ def __init__(
+ self,
+ certs: Mapping[str, "Certificate"],
+ root_dir: Path,
+ name: str = "dataset name",
+ description: str = "dataset_description",
+ ):
self._root_dir = root_dir
self.timestamp = datetime.now()
- self.sha256_digest = 'not implemented'
+ self.sha256_digest = "not implemented"
self.name = name
self.description = description
self.certs = certs
@@ -47,27 +51,27 @@ class Dataset(ABC):
@property
def web_dir(self) -> Path:
- return self.root_dir / 'web'
+ return self.root_dir / "web"
@property
def auxillary_datasets_dir(self) -> Path:
- return self.root_dir / 'auxillary_datasets'
+ return self.root_dir / "auxillary_datasets"
@property
def cpe_dataset_path(self) -> Path:
- return self.auxillary_datasets_dir / 'cpe_dataset.json'
+ return self.auxillary_datasets_dir / "cpe_dataset.json"
@property
def cve_dataset_path(self) -> Path:
- return self.auxillary_datasets_dir / 'cve_dataset.json'
+ return self.auxillary_datasets_dir / "cve_dataset.json"
@property
def nist_cve_cpe_matching_dset_path(self) -> Path:
- return self.auxillary_datasets_dir / 'nvdcpematch-1.0.json'
+ return self.auxillary_datasets_dir / "nvdcpematch-1.0.json"
@property
def json_path(self) -> Path:
- return self.root_dir / (self.name + '.json')
+ return self.root_dir / (self.name + ".json")
def __contains__(self, item):
if not issubclass(type(item), Certificate):
@@ -80,8 +84,8 @@ class Dataset(ABC):
def __getitem__(self, item: str):
return self.certs.__getitem__(item.lower())
- def __setitem__(self, key: str, value: 'Certificate'):
- self.certs.__setitem__(key.lower(), value) # type: ignore
+ def __setitem__(self, key: str, value: "Certificate"):
+ self.certs.__setitem__(key.lower(), value) # type: ignore
def __len__(self) -> int:
return len(self.certs)
@@ -92,65 +96,70 @@ class Dataset(ABC):
return self.certs == other.certs
def __str__(self) -> str:
- return str(type(self).__name__) + ':' + self.name + ', ' + str(len(self)) + ' certificates'
+ return str(type(self).__name__) + ":" + self.name + ", " + str(len(self)) + " certificates"
def to_dict(self):
- return {'timestamp': self.timestamp, 'sha256_digest': self.sha256_digest,
- 'name': self.name, 'description': self.description,
- 'n_certs': len(self), 'certs': list(self.certs.values())}
+ return {
+ "timestamp": self.timestamp,
+ "sha256_digest": self.sha256_digest,
+ "name": self.name,
+ "description": self.description,
+ "n_certs": len(self),
+ "certs": list(self.certs.values()),
+ }
@classmethod
def from_dict(cls, dct: Dict):
- certs = {x.dgst: x for x in dct['certs']}
- dset = cls(certs, Path('../'), dct['name'], dct['description'])
- if len(dset) != (claimed := dct['n_certs']):
+ certs = {x.dgst: x for x in dct["certs"]}
+ dset = cls(certs, Path("../"), dct["name"], dct["description"])
+ if len(dset) != (claimed := dct["n_certs"]):
logger.error(
- f'The actual number of certs in dataset ({len(dset)}) does not match the claimed number ({claimed}).')
+ f"The actual number of certs in dataset ({len(dset)}) does not match the claimed number ({claimed})."
+ )
return dset
@classmethod
- def from_json(cls, input_path: Union[str, Path]):
+ def from_json(cls: Type[T], input_path: Union[str, Path]) -> T:
dset = ComplexSerializableType.from_json(input_path)
dset.root_dir = Path(input_path).parent.absolute()
dset.set_local_paths()
return dset
def set_local_paths(self):
- raise NotImplementedError('Not meant to be implemented by the base class.')
+ raise NotImplementedError("Not meant to be implemented by the base class.")
@abstractmethod
def get_certs_from_web(self):
- raise NotImplementedError('Not meant to be implemented by the base class.')
+ raise NotImplementedError("Not meant to be implemented by the base class.")
@abstractmethod
def convert_all_pdfs(self):
- raise NotImplementedError('Not meant to be implemented by the base class.')
+ raise NotImplementedError("Not meant to be implemented by the base class.")
@abstractmethod
def download_all_pdfs(self, cert_ids: Optional[Set[str]] = None):
- raise NotImplementedError('Not meant to be implemented by the base class.')
+ raise NotImplementedError("Not meant to be implemented by the base class.")
@staticmethod
def _download_parallel(urls: Collection[str], paths: Collection[Path], prune_corrupted: bool = True):
- exit_codes = cert_processing.process_parallel(helpers.download_file,
- list(zip(urls, paths)),
- config.n_threads,
- unpack=True)
+ exit_codes = cert_processing.process_parallel(
+ helpers.download_file, list(zip(urls, paths)), config.n_threads, unpack=True
+ )
n_successful = len([e for e in exit_codes if e == requests.codes.ok])
- logger.info(f'Successfully downloaded {n_successful} files, {len(exit_codes) - n_successful} failed.')
+ logger.info(f"Successfully downloaded {n_successful} files, {len(exit_codes) - n_successful} failed.")
for url, e in zip(urls, exit_codes):
if e != requests.codes.ok:
- logger.error(f'Failed to download {url}, exit code: {e}')
+ logger.error(f"Failed to download {url}, exit code: {e}")
if prune_corrupted is True:
for p in paths:
if p.exists() and p.stat().st_size < constants.MIN_CORRECT_CERT_SIZE:
- logger.error(f'Corrupted file at: {p}')
+ logger.error(f"Corrupted file at: {p}")
p.unlink()
def _prepare_cpe_dataset(self, download_fresh_cpes: bool = False):
- logger.info('Preparing CPE dataset.')
+ logger.info("Preparing CPE dataset.")
if not self.auxillary_datasets_dir.exists():
self.auxillary_datasets_dir.mkdir(parents=True)
@@ -162,8 +171,10 @@ class Dataset(ABC):
return cpe_dataset
- def _prepare_cve_dataset(self, download_fresh_cves: bool = False, use_nist_cpe_matching_dict: bool = True) -> CVEDataset:
- logger.info('Preparing CVE dataset.')
+ def _prepare_cve_dataset(
+ self, download_fresh_cves: bool = False, use_nist_cpe_matching_dict: bool = True
+ ) -> CVEDataset:
+ logger.info("Preparing CVE dataset.")
if not self.auxillary_datasets_dir.exists():
self.auxillary_datasets_dir.mkdir(parents=True)
@@ -177,7 +188,7 @@ class Dataset(ABC):
return cve_dataset
def _compute_candidate_versions(self):
- logger.info('Computing heuristics: possible product versions in sample name')
+ logger.info("Computing heuristics: possible product versions in sample name")
for cert in self:
cert.compute_heuristics_version()
@@ -186,13 +197,22 @@ class Dataset(ABC):
"""
Filters out very weak CPE matches that don't improve our database.
"""
- if cpe.title and (cpe.version == '-' or cpe.version == '*') and not any(char.isdigit() for char in cpe.title):
+ if (
+ cpe.title
+ and (cpe.version == "-" or cpe.version == "*")
+ and not any(char.isdigit() for char in cpe.title)
+ ):
return False
- elif not cpe.title and cpe.item_name and (cpe.version == '-' or cpe.version == '*') and not any(char.isdigit() for char in cpe.item_name):
+ elif (
+ not cpe.title
+ and cpe.item_name
+ and (cpe.version == "-" or cpe.version == "*")
+ and not any(char.isdigit() for char in cpe.item_name)
+ ):
return False
return True
- logger.info('Computing heuristics: Finding CPE matches for certificates')
+ logger.info("Computing heuristics: Finding CPE matches for certificates")
cpe_dset = self._prepare_cpe_dataset(download_fresh_cpes)
if not cpe_dset.was_enhanced_with_vuln_cpes:
cve_dset = self._prepare_cve_dataset(False)
@@ -201,7 +221,7 @@ class Dataset(ABC):
clf = CPEClassifier(config.cpe_matching_threshold, config.cpe_n_max_matches)
clf.fit([x for x in cpe_dset if filter_condition(x)])
- for cert in helpers.tqdm(self, desc='Predicting CPE matches with the classifier'):
+ for cert in helpers.tqdm(self, desc="Predicting CPE matches with the classifier"):
cert.compute_heuristics_cpe_match(clf)
return clf, cpe_dset
@@ -214,51 +234,52 @@ class Dataset(ABC):
def to_label_studio_json(self, output_path: Union[str, Path]):
lst = []
for cert in [x for x in self if x.heuristics.cpe_matches]:
- dct = {'text': cert.label_studio_title}
+ dct = {"text": cert.label_studio_title}
candidates = [x[1].title for x in cert.heuristics.cpe_matches]
- candidates += ['No good match'] * (config.cc_cpe_max_matches - len(candidates))
- options = ['option_' + str(x) for x in range(1, 21)]
+ candidates += ["No good match"] * (config.cc_cpe_max_matches - len(candidates))
+ options = ["option_" + str(x) for x in range(1, 21)]
dct.update({o: c for o, c in zip(options, candidates)})
lst.append(dct)
- with Path(output_path).open('w') as handle:
+ with Path(output_path).open("w") as handle:
json.dump(lst, handle, indent=4)
@serialize
def load_label_studio_labels(self, input_path: Union[str, Path]):
- with Path(input_path).open('r') as handle:
+ with Path(input_path).open("r") as handle:
data = json.load(handle)
cpe_dset = self._prepare_cpe_dataset()
- logger.info('Translating label studio matches into their CPE representations and assigning to certificates.')
- for annotation in helpers.tqdm([x for x in data if 'verified_cpe_match' in x], desc='Translating label studio matches'):
- match_keys = annotation['verified_cpe_match']
- match_keys = [match_keys] if isinstance(match_keys, str) else match_keys['choices']
- match_keys = [x.lstrip('$') for x in match_keys]
- predicted_annotations = [annotation[x] for x in match_keys if annotation[x] != 'No good match']
+ logger.info("Translating label studio matches into their CPE representations and assigning to certificates.")
+ for annotation in helpers.tqdm(
+ [x for x in data if "verified_cpe_match" in x], desc="Translating label studio matches"
+ ):
+ match_keys = annotation["verified_cpe_match"]
+ match_keys = [match_keys] if isinstance(match_keys, str) else match_keys["choices"]
+ match_keys = [x.lstrip("$") for x in match_keys]
+ predicted_annotations = [annotation[x] for x in match_keys if annotation[x] != "No good match"]
cpes: Set[Optional[CPE]] = set()
for x in predicted_annotations:
if x not in cpe_dset.title_to_cpes:
- print(f'Error: {x} not in dataset')
+ print(f"Error: {x} not in dataset")
else:
to_update = cpe_dset.title_to_cpes[x]
if to_update and not cpes:
cpes = to_update
elif to_update and cpes:
- # TODO: This was here like cpes = cpes.update(to_update), but update() does not return anything.
- # Did you try to hack something using that or was that just a typo?
+ # TODO: This was here like cpes = cpes.update(to_update), but update() does not return anything.
+ # Did you try to hack something using that or was that just a typo?
cpes.update(to_update)
-
# cpes = set(itertools.chain.from_iterable([cpe_dset.title_to_cpes.get(x, []) for x in predicted_annotations]))
# distinguish between FIPS and CC
- if '\n' in annotation['text']:
- cert_name = annotation['text'].split('\nModule name: ')[1].split('\n')[0]
+ if "\n" in annotation["text"]:
+ cert_name = annotation["text"].split("\nModule name: ")[1].split("\n")[0]
else:
- cert_name = annotation['text']
+ cert_name = annotation["text"]
certs = self.get_certs_from_name(cert_name)
@@ -266,7 +287,7 @@ class Dataset(ABC):
c.heuristics.verified_cpe_matches = {x.uri for x in cpes if x is not None} if cpes else None
def get_certs_from_name(self, name: str) -> List[Certificate]:
- raise NotImplementedError('Not meant to be implemented by the base class.')
+ raise NotImplementedError("Not meant to be implemented by the base class.")
def enrich_automated_cpes_with_manual_labels(self):
"""
@@ -276,27 +297,32 @@ class Dataset(ABC):
if not cert.heuristics.cpe_matches and cert.heuristics.verified_cpe_matches:
cert.heuristics.cpe_matches = cert.heuristics.verified_cpe_matches
elif cert.heuristics.cpe_matches and cert.heuristics.verified_cpe_matches:
- cert.heuristics.cpe_matches = set(cert.heuristics.cpe_matches).union(set(cert.heuristics.verified_cpe_matches))
+ cert.heuristics.cpe_matches = set(cert.heuristics.cpe_matches).union(
+ set(cert.heuristics.verified_cpe_matches)
+ )
@serialize
def compute_related_cves(self, download_fresh_cves: bool = False, use_nist_cpe_matching_dict: bool = True):
- logger.info('Retrieving related CVEs to verified CPE matches')
+ logger.info("Retrieving related CVEs to verified CPE matches")
cve_dset = self._prepare_cve_dataset(download_fresh_cves, use_nist_cpe_matching_dict)
self.enrich_automated_cpes_with_manual_labels()
cpe_rich_certs = [x for x in self if x.heuristics.cpe_matches]
if not cpe_rich_certs:
- logger.error('No certificates with verified CPE match detected. You must run dset.manually_verify_cpe_matches() first. Returning.')
+ logger.error(
+ "No certificates with verified CPE match detected. You must run dset.manually_verify_cpe_matches() first. Returning."
+ )
return
relevant_cpes = set(itertools.chain.from_iterable([x.heuristics.cpe_matches for x in cpe_rich_certs]))
cve_dset.filter_related_cpes(relevant_cpes)
- for cert in helpers.tqdm(cpe_rich_certs, desc='Computing related CVES'):
+ for cert in helpers.tqdm(cpe_rich_certs, desc="Computing related CVES"):
cert.compute_heuristics_related_cves(cve_dset)
n_vulnerable = len([x for x in cpe_rich_certs if x.heuristics.related_cves])
- n_vulnerabilities = sum(
- [len(x.heuristics.related_cves) for x in cpe_rich_certs if x.heuristics.related_cves])
- logger.info(f'In total, we identified {n_vulnerabilities} vulnerabilities in {n_vulnerable} vulnerable certificates.')
+ n_vulnerabilities = sum([len(x.heuristics.related_cves) for x in cpe_rich_certs if x.heuristics.related_cves])
+ logger.info(
+ f"In total, we identified {n_vulnerabilities} vulnerabilities in {n_vulnerable} vulnerable certificates."
+ )